文章详情

短信预约-IT技能 免费直播动态提醒

请输入下面的图形验证码

提交验证

短信预约提醒成功

python实现最大熵模型

2023-01-31 01:56

关注

# encoding: utf-8
'''
Created on 2017-8-7
根据李航<<统计学习方法>>实现
'''

from collections import defaultdict
import math

class MaxEnt(object):
    def __init__(self):
        self.feats = defaultdict(int)
        self.trainset = []
        self.labels = set()  
      
    def load_data(self, file):
        for line in open(file):
            fields = line.strip().split()
            
            # 数据共3列。第一列为标签,二三列为特征
            if len(fields) < 2: continue
            label = fields[0]
            self.labels.add(label)
            for f in set(fields[1:]):
                # (label,f) tuple is feature 
                self.feats[(label, f)] += 1
            self.trainset.append(fields)
            
    def _initparams(self):
        self.size = len(self.trainset)
        
        self.M = max([len(record) - 1 for record in self.trainset]) # P91中的M
        
        # 计算P82页最下面的期望
        self.ep_ = [0.0] * len(self.feats)  # 保存期望值
        for i, f in enumerate(self.feats):
            self.ep_[i] = float(self.feats[f]) / float(self.size)
            # each feature function correspond to id
            self.feats[f] = i

        # 初始化需要学习的参数的值
        self.w = [0.0] * len(self.feats)
        self.lastw = self.w
        
        
    def probwgt(self, features, label):
        '''
                        辅助函数:计算P85中的公式6.22中的分子
        '''
        wgt = 0.0
        for f in features:
            print (self.feats[(label, f)])
            if (label, f)in self.feats:
                wgt += self.w[self.feats[(label, f)]]
        return math.exp(wgt)


    
    def calprob(self, features):
        '''
                        计算P85中的公式6.22的条件概率P(y|x)
        '''
        wgts = [(self.probwgt(features, label), label) for label in self.labels]
        Z = sum([ w for w, label in wgts])
        prob = [ (w / Z, label) for w, label in wgts]
        return prob 
    
                       
    def Ep(self):
        '''
                        计算P83页最上面的期望
        '''
        eps = [0.0] * len(self.feats)
        for record in self.trainset:
            features = record[1:]
            
            # 计算 p(y|x)
            probs = self.calprob(features)
            for f in features:
                for prob, label in probs:
                    if (label, f) in self.feats:     # only focus on features from training data.
                        idx = self.feats[(label, f)]
                        eps[idx] += prob * (1.0 / self.size) # 计算期望 sum(P(x) * P(y|x) * f(x,y))。 其中P(x) = 1 / N
        return eps
    
    def _convergence(self, lastw, w):
        for w1, w2 in zip(lastw, w):
            if abs(w1 - w2) >= 0.01:
                return False
        return True
                
    def train(self, max_iter=1000):
        self._initparams()
        for i in range(max_iter):
            print ('iter %d ...' % (i + 1))
            self.ep = self.Ep()           
            self.lastw = self.w[:]  
            for i, w in enumerate(self.w):
                delta = 1.0 / self.M * math.log(self.ep_[i] / self.ep[i])   # P91 公式6.34
                self.w[i] += delta
            
            # 是否满足收敛条件    
            if self._convergence(self.lastw, self.w):
                break

            
    def predict(self, input):
        features = input.strip().split()
        prob = self.calprob(features)
        prob.sort(reverse=True)
        return prob 

if __name__ == "__main__":
    maxent = MaxEnt()
    maxent.load_data("input.data")
    maxent.train(100)
    prob = maxent.predict("Sunny  Sad")
    print (prob)


github上发现的一份最大熵模型实现代码。具体链接找不到了。


阅读原文内容投诉

免责声明:

① 本站未注明“稿件来源”的信息均来自网络整理。其文字、图片和音视频稿件的所属权归原作者所有。本站收集整理出于非商业性的教育和科研之目的,并不意味着本站赞同其观点或证实其内容的真实性。仅作为临时的测试数据,供内部测试之用。本站并未授权任何人以任何方式主动获取本站任何信息。

② 本站未注明“稿件来源”的临时测试数据将在测试完成后最终做删除处理。有问题或投稿请发送至: 邮箱/279061341@qq.com QQ/279061341

软考中级精品资料免费领

  • 历年真题答案解析
  • 备考技巧名师总结
  • 高频考点精准押题
  • 2024年上半年信息系统项目管理师第二批次真题及答案解析(完整版)

    难度     813人已做
    查看
  • 【考后总结】2024年5月26日信息系统项目管理师第2批次考情分析

    难度     354人已做
    查看
  • 【考后总结】2024年5月25日信息系统项目管理师第1批次考情分析

    难度     318人已做
    查看
  • 2024年上半年软考高项第一、二批次真题考点汇总(完整版)

    难度     435人已做
    查看
  • 2024年上半年系统架构设计师考试综合知识真题

    难度     224人已做
    查看

相关文章

发现更多好内容

猜你喜欢

AI推送时光机
位置:首页-资讯-后端开发
咦!没有更多了?去看看其它编程学习网 内容吧
首页课程
资料下载
问答资讯