目录

一、引言

二、贝叶斯网络

三、最大后验估计 (Maximum A Posteriori Estimation, MAP)

四、马尔可夫网络 (Markov Network / Markov Random Field, MRF)

五、朴素贝叶斯文本分类与马尔可夫网络图像去噪的算法实现步骤与Python代码实现

六、总结


一、引言

在本文中,我们将讨论无监督学习中的数据分布建模问题。当我们需要在一个数据集上完成某个任务时,数据集中的样本分布显然是最基本的要素。面对不同的数据分布,我们可能针对同一任务采用完全不同的算法。例如,如果样本有明显的线性相关关系,我们就可以考虑用基于线性模型的算法解决问题;如果样本呈高斯分布,我们可能会使用高斯分布的各种性质来简化任务的要求。如果在数据分布不明显的情况下,我们可能会使用数据降维算法(如主成分分析算法)来尽可能提取出数据的关键特征。因此,如何建模数据集中样本关于其各个特征的分布,就成了一个相当关键的问题。

在生活中,我们常常会看到这样的情况:秋天时,一群人在路上行走,黄叶遍地,风光无限。但是,在同一个温度下,不同人的衣着也有差异。有人穿了厚厚的大衣,有人穿了长袖长裤,还有人穿着短袖或者裙子。可以设想,如果气温再高一些,穿短袖或者裙子的人会增加;如果气温再低一些,恐怕就不会有人再穿短袖或者裙子了。因此我们可以认为,人群穿衣选择的概率分布受到天气的影响。

我们先从最简单的表格数据看起,假设表中是天气和人群中衣服选择的部分数据,我们可以直接从中写出最简单的数据分布:

                           P(天气 = 热,衣服 = 衬衫) = 0.48

天气 衣服 概率
衬衫 0.48
大衣 0.12
衬衫 0.08
大衣 0.32

以此类推,我们可以把表中的每一行都写出来,得到样本的分布。但是,这样的做法显然过于低效。当特征的数目增加时,我们按此建模的复杂度将呈指数增长。因此,我们需要设法寻找不同特征之间的相关性,降低模型的复杂度。例如,根据生活常识,我们可以认为人们选择衣服的概率应该是和天气有关的。天气热时,人们更倾向于选择衬衫,而天气冷时倾向于大衣。这样,我们可以将上面的分布转化为条件概率:

              P(天气 = 热) = 0.6,P(衣服 = 衬衫|天气 = 热) = 0.8

通过这种方式,我们可以建立起样本不同特征之间的关系。如果用随机变量t表示天气,c表示衣服,那么上述的关系可以表示为图中的结构。

                             

在图中,从t指向c的箭头表示随机变量c依赖于t。把样本的所有特征按依赖关系列出来,每个特征作为一个顶点,每对依赖关系作为一条边,就形成了一张概率图。我们可以通过概率图中体现的不同特征之间的关系,推断出数据的概率分布。由依赖关系构成的图是有向图,称为贝叶斯网络。如果我们知道两个特征之间相关,但没有明确的单向依赖关系,就可以用无向图来建模,称为马尔可夫网络。下面,我们就来介绍这两种概率图模型的具体内容及Python代码实战。

二、贝叶斯网络

1. 核心思想与结构

朴素贝叶斯网络是一种基于贝叶斯定理的简单概率分类器,其核心假设是:给定类别标签时,所有特征之间是条件独立的。

2. 数学公式详解

(1)贝叶斯定理 (Bayes' Theorem):

对于类别 C_k 和特征集 X = (X_1, X_2, ..., X_n),贝叶斯定理表述为:

                                             P(C_k | X) = \frac{P(X | C_k) \cdot P(C_k)}{P(X)}

其中:

P(C_k | X):后验概率 (Posterior Probability),即在已知特征 X 的条件下,样本属于类别 C_k 的概率。

P(X | C_k):似然性 (Likelihood),即在样本属于类别 C_k 的条件下,观察到特征 X 的概率。

P(C_k):先验概率 (Prior Probability),即类别 C_k 自身发生的概率,与特征无关。

P(X):证据 (Evidence),即特征 X 发生的概率,也称为边缘似然。对于所有类别,P(X) 是相同的,因此在比较不同类别的后验概率时可以忽略。

(2)朴素的特征条件独立性假设:

这是朴素贝叶斯的核心。它假设在给定类别 C_k 的条件下,各个特征 X_i 是相互独立的。因此,似然性 P(X | C_k) 可以被分解为:

                         P(X | C_k) = P(X_1, X_2, ..., X_n | C_k) = \prod_{i=1}^{n} P(X_i | C_k)

(3)后验概率计算:

结合贝叶斯定理和独立性假设,后验概率的计算变为:

                                            P(C_k | X) = \frac{P(C_k) \prod_{i=1}^{n} P(X_i | C_k)}{P(X)}

由于 P(X) 是归一化因子,在分类决策时,我们通常关注分子部分,即:

                                        P(C_k | X) \propto P(C_k) \prod_{i=1}^{n} P(X_i | C_k)

(4)分类决策规则:

朴素贝叶斯分类器选择具有最大后验概率的类别作为预测结果:

              \hat{C} = \arg\max_{C_k} P(C_k | X) = \arg\max_{C_k} \left( P(C_k) \prod_{i=1}^{n} P(X_i | C_k) \right)

为了避免多个小概率相乘导致的浮点数下溢问题,通常使用对数概率进行计算:

                        \hat{C} = \arg\max_{C_k} \left( \log P(C_k) + \sum_{i=1}^{n} \log P(X_i | C_k) \right)

(5)参数估计 (Parameter Estimation):

模型中的概率需要从训练数据中估计。

a.先验概率 P(C_k): 使用最大似然估计 (MLE)

                       P(C_k) =属于类别C_k的样本数量/总样本数量 = \frac{N_k}{N}

b.  条件概率 P(X_i | C_k):

对于离散特征 (如文本分类中的词语):

    P(X_i=x_{ij} | C_k) =类别 C_k 中特征X_i取值为x_{ij}的样本数量/类别C_k的总样本数量= \frac{N_{k,ij}}{N_k}

其中 x_{ij} 是特征 X_i 的第 j 个可能取值。

对于连续特征 (如高斯朴素贝叶斯):假设特征 X_i 在类别 C_k 下服从高斯分布 N(\mu_{ik}, \sigma_{ik}^2)。则:

                                   P(X_i=x_i | C_k) = \frac{1}{\sqrt{2\pi}\sigma_{ik}} \exp\left(-\frac{(x_i - \mu_{ik})^2}{2\sigma_{ik}^2}\right)

其中均值 \mu_{ik} 和方差 \sigma_{ik}^2 从训练数据中属于类别 C_k 的样本的特征 X_i 值估计得到。

平滑技术 (Smoothing - 解决零概率问题):

如果某个特征值在训练集中某个类别下从未出现,其条件概率会为0,导致整个后验概率为0。拉普拉斯平滑 (或加法平滑) 是一种常用方法。对于离散特征:

                                                   P(X_i=x_{ij} | C_k) = \frac{N_{k,ij} + \alpha}{N_k + \alpha \cdot V_i}

    其中:

    \alpha 是平滑参数 (通常 \alpha=1,称为加一平滑)。

    V_i 是特征 X_i 可能取值的数量。

3.优点

简单高效: 算法逻辑简单,计算速度快,易于实现。

对小数据集表现良好:即使训练数据量不大,也能获得不错的效果,因为它只需要估计先验概率和条件概率。

处理高维数据:在特征数量很多的情况下(如文本分类中的词汇表大小),依然能有效工作。

对无关特征不敏感: 如果数据中存在与分类任务无关的特征,朴素贝叶斯受到的影响相对较小。

4. 缺点

特征条件独立性假设过于理想化: 这是其“朴素”的来源。在现实世界中,特征之间往往存在一定的关联性。这个强假设可能会限制模型的精度,尤其是在特征高度相关的情况下。

零概率问题: 如果在训练阶段某个特征值在某个类别下从未出现过,那么整个后验概率为0。这通常通过平滑技术(如拉普拉斯平滑/加法平滑)来解决,即给所有计数都加上一个小的正数。

5. 常见变体

根据特征的数据类型不同,朴素贝叶斯有几种常见变体:

高斯朴素贝叶斯 (Gaussian Naive Bayes): 适用于连续型特征,假设特征服从高斯分布。

多项式朴素贝叶斯 (Multinomial Naive Bayes):常用于文本分类,特征通常是词频计数。

伯努利朴素贝叶斯 (Bernoulli Naive Bayes):也用于文本分类,特征是二值的(表示词语是否出现)。

6. 应用

广泛应用于文本分类(垃圾邮件过滤、情感分析、新闻分类)、医疗诊断、推荐系统等。

三、最大后验估计 (Maximum A Posteriori Estimation, MAP)

1. 定义与目的

最大后验估计 (MAP) 是一种在贝叶斯统计中估计未知参数 \theta 的方法。它结合了观测数据 D 和关于参数 \theta 的先验知识 P(\theta),目标是找到使参数的后验概率 P(\theta|D) 最大化的那个参数值。

2. 数学公式详解

后验概率 (Posterior Probability):

根据贝叶斯定理,参数 \theta 的后验概率为:

                                                   P(\theta | D) = \frac{P(D | \theta) \cdot P(\theta)}{P(D)}

    其中:

   P(D | \theta):似然函数 (Likelihood),给定参数 \theta 时观测到数据 D 的概率。

   P(\theta):先验概率 (Prior Probability),在观测数据之前我们对参数 \theta 的信念或知识。

   P(D):证据 (Evidence) 或边缘似然,P(D) = \int P(D|\theta)P(\theta) d\theta (对于连续参数) 或 P(D) = \sum_{\theta} P(D|\theta)P(\theta) (对于离散参数)。它是归一化常数。

MAP 估计的定义:

MAP 估计 \hat{\theta}_{\text{MAP}} 是使后验概率 P(\theta | D) 最大化的参数值:

                                                 \hat{\theta}_{\text{MAP}} = \arg\max_{\theta} P(\theta | D)

由于 P(D) 与 \theta 无关,最大化 P(\theta | D) 等价于最大化分子 P(D | \theta) P(\theta)

                                           \hat{\theta}_{\text{MAP}} = \arg\max_{\theta} [P(D | \theta) \cdot P(\theta)]

与最大似然估计 (MLE) 的关系:

最大似然估计 \hat{\theta}_{\text{MLE}} 旨在最大化似然函数:

                                                  \hat{\theta}_{\text{MLE}} = \arg\max_{\theta} P(D | \theta)

如果先验分布 P(\theta) 是一个均匀分布 (即对所有 \theta 值都赋予相同的概率,表示没有特别的先验偏好),那么 MAP 估计就退化为 MLE 估计。

对数后验概率:

为了计算方便,通常最大化对数后验概率:

                                    \hat{\theta}_{\text{MAP}} = \arg\max_{\theta} [\log P(D | \theta) + \log P(\theta)]

这里,\log P(D | \theta) 是对数似然,\log P(\theta) 是对数先验。对数先验项可以看作是对模型的正则化项。

示例:朴素贝叶斯中的拉普拉斯平滑作为MAP估计

在估计朴素贝叶斯中的条件概率 P(X_i=x_{ij} | C_k) 时,如果使用狄利克雷分布 (Dirichlet distribution) 作为多项式分布参数的共轭先验,那么对这些参数进行MAP估计的结果就等价于进行了拉普拉斯平滑或更一般的加法平滑。

例如,对于一个具有 V 个可能值的特征,其在类别 C_k 下的参数(即每个值出现的概率)\vec{p}_k = (p_{k1}, ..., p_{kV}) 服从多项式分布。如果给 \vec{p}_k 一个对称的狄利克雷先验 Dir(\alpha, ..., \alpha),那么 \log P(\vec{p}_k) = (\alpha-1) \sum_j \log p_{kj} + \text{const}

结合对数似然 \log P(D_k | \vec{p}_k) = \sum_j N_{kj} \log p_{kj} (其中 D_k 是类别 C_k 下的数据,N_{kj} 是特征值 j 在类别 C_k 中出现的次数),最大化 \sum_j N_{kj} \log p_{kj} + (\alpha-1) \sum_j \log p_{kj} 会得到 p_{kj} = \frac{N_{kj} + \alpha - 1}{\sum_l (N_{kl} + \alpha - 1)}。如果令平滑参数为 \alpha' (通常的拉普拉斯平滑参数),则对应于狄利克雷先验中的 \alpha = \alpha' + 1

3.作用与优势

正则化/防止过拟合:先验分布 P(θ) 可以起到正则化项的作用。例如,如果先验倾向于选择较小的参数值,那么即使数据本身可能导致参数值很大(可能过拟合),MAP也会将其拉向一个更“合理”的范围。这在数据量较少时尤其有用。

引入先验知识:允许分析者将关于参数的已有信念或领域知识整合到估计过程中。

提供点估计:像MLE一样,MAP也提供参数的一个具体估计值,而不是一个完整的后验分布(像完全贝叶斯推断那样)。

四、马尔可夫网络 (Markov Network / Markov Random Field, MRF)

1. 定义与结构

马尔可夫网络是一种使用无向图 G=(V, E) 来表示一组随机变量 X = (X_1, ..., X_n) 之间联合概率分布的概率图模型。

节点 v \in V 代表随机变量 X_v

(u,v) \in E 表示变量 X_u 和 X_v 之间存在直接的概率依赖关系。

2. 核心特性 (马尔可夫性质)

MRF满足以下条件独立性假设:

局部马尔可夫性 (Local Markov Property): 给定一个节点的所有直接邻居节点的值,该节点条件独立于图中所有其他非邻居节点。

                                               X_v \perp X_{V \setminus (\{v\} \cup N(v))} | X_{N(v)}

    其中 N(v) 是节点 v 的邻居集合。

成对马尔可夫性 (Pairwise Markov Property): 对于图中任意两个没有直接边相连的非邻居节点 X_uX_v,在给定图中所有其他节点的条件下,这两个节点是条件独立的。

                                           X_u \perp X_v | X_{V \setminus \{u,v\}} \quad \text{if } (u,v) \notin E

全局马尔可夫性 (Global Markov Property): 如果节点集A和节点集B被节点集C在图中分离,那么在给定节点集C的值的条件下,节点集A和节点集B是条件独立的。

                                                                X_A \perp X_B | X_C

3. 因子分解与势函数 (Hammersley-Clifford Theorem)

Hammersley-Clifford定理指出,一个随机场的概率分布 P(X) 是一个关于图 G 的马尔可夫随机场,当且仅当它可以被因子分解为定义在图 G 中所有(最大)团 (cliques) 上的非负势函数 (potential functions) \phi_C(X_C) 的乘积,然后进行归一化:

                                                        P(X) = \frac{1}{Z} \prod_{C \in \mathcal{C}(G)} \phi_C(X_C)

其中:

X_C:表示团 C 中所有变量的一个配置。

\mathcal{C}(G):图 G 中所有(通常是最大)团的集合。

\phi_C(X_C) \ge 0:定义在团 C 上的势函数,衡量团内变量配置的“相容性”或“偏好度”。值越大,该配置越可能。

Z:归一化常数,称为配分函数 (partition function),确保所有可能配置的概率之和为1。

                                                     Z = \sum_{X} \prod_{C \in \mathcal{C}(G)} \phi_C(X_C)

    计算 Z 通常是NP难的,这是MRF推断和学习中的主要挑战。

4. 对数线性模型与吉布斯分布 (Log-Linear Models and Gibbs Distribution)

势函数通常采用严格正的指数形式,即对数线性模型:

                                           \phi_C(X_C) = \exp\left( \sum_{k \in I_C} w_k f_k(X_C) \right)

或者更一般地,直接定义为能量项的指数:

                                                 \phi_C(X_C) = \exp\left( -E_C(X_C) \right)

其中 E_C(X_C) 是与团 C 相关的能量。

将此代入联合概率分布,得到吉布斯分布形式:

                                  P(X) = \frac{1}{Z} \exp\left( \sum_{C \in \mathcal{C}(G)} \sum_{k \in I_C} w_k f_k(X_C) \right)

或者,如果定义总能量函数 E(X) = \sum_{C \in \mathcal{C}(G)} E_C(X_C) (注意这里的能量定义与上面的权重特征形式可能不同,但最终都是求和):

                                                     P(X) = \frac{1}{Z} \exp\left( -E(X) \right)

能量越低,该配置 \(X\) 出现的概率越大。

5.与贝叶斯网络的区别

图结构:贝叶斯网络是有向无环图 (DAG),边表示因果关系或条件依赖。马尔可夫网络是无向图,边表示变量间的对称关系或软约束。

概率分解:贝叶斯网络的联合概率通过条件概率链式法则分解 。马尔可夫网络通过团上的势函数分解。

表达能力:马尔可夫网络可以自然地表示贝叶斯网络难以直接表示的循环依赖关系(例如,图像中相邻像素的相互影响)。

6.推断 (Inference)

推断任务包括计算边缘概率或条件概率。

精确推断在一般MRF中是NP难的,因为需要计算配分函数Z。

常用的近似推断算法包括:

置信传播 (Belief Propagation, BP)及其变体(如Loopy Belief Propagation,用于有环图)。

变分推断 (Variational Inference):用一个更简单的分布来近似真实的后验分布。

马尔可夫链蒙特卡洛 (MCMC) 方法:如吉布斯采样 (Gibbs Sampling),通过从目标分布中采样来近似期望值。

7. 学习 (Learning)

参数学习:给定图结构,从数据中学习势函数的参数。由于配分函数Z的存在,直接的最大似然估计很困难。常用的方法有:

伪似然估计 (Pseudo-Likelihood Estimation)。

对比散度 (Contrastive Divergence, CD),常用于训练受限玻尔兹曼机 (RBM) 等能量模型。

最大间隔马尔可夫网络 (Max-Margin Markov Networks)。

结构学习:从数据中学习图的结构(即哪些变量之间应该有边)。这是一个比参数学习更具挑战性的问题。

8. 应用

马尔可夫网络在许多领域都有重要应用,特别是在需要对变量间的局部依赖关系和上下文信息进行建模的场景:

计算机视觉:图像去噪、图像分割、图像恢复、物体识别、立体视觉。

自然语言处理:词性标注、命名实体识别、信息抽取(通常以条件随机场CRF的形式出现,CRF是MRF的一种判别式模型)。

生物信息学:基因序列分析、蛋白质结构预测。

社交网络分析:分析网络结构和影响力传播。

统计物理:如伊辛模型,用于描述粒子系统。

五、朴素贝叶斯文本分类与马尔可夫网络图像去噪的算法实现步骤与Python代码实现

(一)算法实现步骤

朴素贝叶斯文本分类的核心数学公式

1.  贝叶斯定理 (Bayes' Theorem):

                                                   P(A|B) = \frac{P(B|A) \cdot P(A)}{P(B)}

    其中:

    P(A|B):后验概率

    P(B|A):似然性

    P(A):先验概率

    P(B):证据

2.  朴素贝叶斯假设下的文本概率 (Conditional Probability of Text given Class with Naive Assumption):

    如果文本包含词语 w_1, w_2, \dots, w_n,在类别 c 下,该文本出现的概率为:

                        P(\text{text}|c) = P(w_1, w_2, \dots, w_n|c) = \prod_{i=1}^{n} P(w_i|c)

3.  文本分类的后验概率 (Posterior Probability for Text Classification):

                                           P(c|\text{text}) = \frac{P(\text{text}|c) \cdot P(c)}{P(\text{text})}

    由于 P(\text{text}) 对所有类别相同,通常关注:

                                                 P(c|\text{text}) \propto P(\text{text}|c) \cdot P(c)

    结合朴素假设:

                                          P(c|\text{text}) \propto \left( \prod_{i=1}^{n} P(w_i|c) \right) \cdot P(c)

4.  分类决策规则 (Classification Decision Rule):

    选择使后验概率最大化的类别 c^*:

                                    c^* = \arg\max_{c} \left( \prod_{i=1}^{n} P(w_i|c) \right) \cdot P(c)

    使用对数概率以避免下溢:

                          c^* = \arg\max_{c} \left( \sum_{i=1}^{n} \log P(w_i|c) + \log P(c) \right)

5.  参数估计 - 先验概率 P(c) (Prior Probability Estimation):

                    P(c) = 文档集中属于类别c的文档数量/文档集总数量 

6.  参数估计 - 条件概率 P(w|c) (Conditional Probability Estimation - MLE):

     P(w|c) = 词 w 在类别 c 的文档中出现的总次数/类别 c 中所有词语的总次数

7.  拉普拉斯平滑 (Laplace Smoothing / Additive Smoothing):

     P(w|c) = (词 w 在类别 c 的文档中出现的总次数 + \alpha )/(类别 c 中所有词语的总次数 + \alpha \cdot V )

    其中:

   \alpha 是平滑参数 (通常\alpha=1)

    V是词汇表的大小

马尔可夫随机场 (MRF) 图像去噪的核心数学公式

1.  局部马尔可夫性质 (Local Markov Property):

    对于任意节点 x_i,令 N(i) 为其邻域:

                                      P(x_i | x_{V \setminus \{i\}}) = P(x_i | x_{N(i)})

2.  吉布斯分布 (Gibbs Distribution):

    MRF的概率分布可以表示为:

                                                        P(x) = \frac{1}{Z} \exp\left(-\frac{E(x)}{T}\right)

    其中:

    x:图像所有像素值的配置

    E(x):能量函数, E(x) = \sum_{C \in \mathcal{C}} \phi_C(x_C)\phi_C(x_C) 是团 C 上的势函数。

    Z:归一化因子(配分函数), Z = \sum_x \exp\left(-\frac{E(x)}{T}\right)

    T:温度参数 (在优化中可设为1)

3.  图像去噪的能量函数设计 (Energy Function for Image Denoising):

    总能量函数  E(x, y) = E_{\text{data}}(x, y) + E_{\text{smooth}}(x)

    保真项 (Data Term / Likelihood Energy):

                                        E_{\text{data}}(x, y) = \sum_{i \in V} (x_i - y_i)^2

       其中 x_i 是恢复图像像素值,y_i 是噪声图像像素值。

        平滑项 (Prior Term / Smoothness Energy):

                                  E_{\text{smooth}}(x) = \sum_{(i,j) \in E} \beta \cdot \psi(x_i, x_j)

        其中 \psi(x_i, x_j) 可以是 (x_i - x_j)^2\mathbb{I}(x_i \neq x_j) (指示函数)。

4.  图像去噪的目标 (Objective for Image Denoising):

    找到使总能量最小的图像 x^*:

                                     x^* = \arg\min_x E(x,y)

   5. 使用的局部能量函数 (Local Energy Function used in the provided code analysis for MRFImageDenoiser):

    这是在ICM算法中,针对单个像素 i (其值为 x_i,对应噪声图像值为 y_i) 及其邻域 N(i) (邻居像素值为 x_j) 计算的局部能量。这种形式常见于二值图像(像素值为 +1 或 -1)的Ising模型。 

                         E_{\text{local}}(x_i, y_i, \{x_j\}_{j \in N(i)}) = -\eta x_i y_i - \beta \sum_{j \in N(i)} x_i x_j

    其中:

    -\eta x_i y_i:数据项(保真项)。如果 x_i 和 y_i 相同,此项为 -\eta,能量较低;如果不同,为 +\eta,能量较高。

   -\beta \sum_{j \in N(i)} x_i x_j:平滑项。如果 x_i 和其邻居 x_j 相同,此项对每个邻居贡献 -\beta,能量较低;如果不同,贡献 +\beta,能量较高。

   \eta:数据项权重。

    \beta:平滑项权重。

ICM算法通过迭代地为每个像素选择能使其上述局部能量最小化的值来更新图像。

(二)Python代码实现

import numpy as np
import matplotlib.pyplot as plt
from matplotlib.backends.backend_tkagg import FigureCanvasTkAgg, NavigationToolbar2Tk
import time
import logging
import tkinter as tk
from tkinter import ttk, filedialog, messagebox, scrolledtext
import threading
from sklearn.feature_extraction.text import TfidfVectorizer, CountVectorizer
from sklearn.naive_bayes import MultinomialNB
from sklearn.metrics import classification_report, confusion_matrix, accuracy_score, precision_recall_fscore_support
from sklearn.model_selection import GridSearchCV, train_test_split
from sklearn.datasets import make_classification
from sklearn.ensemble import RandomForestClassifier, VotingClassifier
from sklearn.svm import SVC
from sklearn.pipeline import Pipeline
from tqdm import tqdm, trange
import seaborn as sns
import argparse
import os
import sys
from multiprocessing import Pool, cpu_count
import warnings
import pickle
from PIL import Image, ImageTk
import io
import datetime

# 配置日志
log_filename = f"ml_application_{datetime.datetime.now().strftime('%Y%m%d_%H%M%S')}.log"
logging.basicConfig(
    level=logging.INFO,
    format='%(asctime)s - %(levelname)s - %(message)s',
    handlers=[
        logging.FileHandler(log_filename),  # 输出到文件
        logging.StreamHandler()  # 输出到控制台
    ]
)
logger = logging.getLogger(__name__)

# 设置中文字体支持
plt.rcParams['font.sans-serif'] = ['SimHei']  # 用来正常显示中文标签
plt.rcParams['axes.unicode_minus'] = False  # 用来正常显示负号

# 忽略不必要的警告
warnings.filterwarnings("ignore", category=UserWarning)
warnings.filterwarnings("ignore", category=FutureWarning)

# 解析命令行参数
parser = argparse.ArgumentParser(description='Python机器学习:朴素贝叶斯文本分类与马尔可夫网络图像去噪实现')
parser.add_argument('--task', type=str, default='both', choices=['nb', 'mrf', 'both'],
                    help='选择运行的任务:nb(朴素贝叶斯)、mrf(马尔可夫网络)或both(两者都运行)')
parser.add_argument('--gui', action='store_true', help='启用图形用户界面')
parser.add_argument('--interactive', action='store_true', help='启用交互模式(命令行)')
parser.add_argument('--categories', type=int, default=4, help='朴素贝叶斯分类中使用的类别数量(2-20)')
parser.add_argument('--alpha', type=float, default=1.0, help='朴素贝叶斯的平滑参数')
parser.add_argument('--noise_rate', type=float, default=0.1, help='图像噪声比例(0-0.5)')
parser.add_argument('--eta', type=float, default=2.0, help='马尔可夫网络的数据项权重')
parser.add_argument('--beta', type=float, default=1.0, help='马尔可夫网络的光滑项权重')
parser.add_argument('--max_iters', type=int, default=10, help='马尔可夫网络最大迭代次数')
parser.add_argument('--save_results', action='store_true', help='保存结果到文件')
parser.add_argument('--output_dir', type=str, default='results', help='结果输出目录')
parser.add_argument('--parallel', action='store_true', help='使用并行处理加速计算')
parser.add_argument('--verbose', action='store_true', help='显示详细输出')
parser.add_argument('--ensemble', action='store_true', help='使用模型集成技术')

args = parser.parse_args()

# 创建输出目录
if args.save_results:
    timestamp = datetime.datetime.now().strftime('%Y%m%d_%H%M%S')
    args.output_dir = f"{args.output_dir}_{timestamp}"
    if not os.path.exists(args.output_dir):
        os.makedirs(args.output_dir)

# 设置随机种子确保结果可复现
np.random.seed(42)


###########################################
# 生成随机数据集
###########################################

class CustomTextDataset:
    """
    生成模拟文本数据集,替代fetch_20newsgroups,避免网络下载失败
    """

    def __init__(self, num_categories=4, samples_per_category=100, vocab_size=1000, random_state=42):
        np.random.seed(random_state)
        self.num_categories = num_categories
        self.samples_per_category = samples_per_category
        self.vocab_size = vocab_size
        self.target_names = [f"category_{i}" for i in range(num_categories)]

        # 生成数据
        self.data, self.target = self._generate_data()

    def _generate_document(self, category, doc_length=100):
        """生成一篇模拟文档"""
        # 为每个类别创建不同的词频分布
        word_probs = np.random.dirichlet(np.ones(self.vocab_size) * (10 if category == 0 else 5))
        # 偏向某些词
        emphasis = np.zeros(self.vocab_size)
        emphasis[category * (self.vocab_size // self.num_categories):(category + 1) * (
                    self.vocab_size // self.num_categories)] = 10
        word_probs = word_probs + emphasis
        word_probs = word_probs / np.sum(word_probs)

        # 从词频分布中采样词
        words = np.random.choice(self.vocab_size, size=doc_length, p=word_probs)
        return " ".join([f"word_{w}" for w in words])

    def _generate_data(self):
        """生成整个数据集"""
        data = []
        target = []

        for category in range(self.num_categories):
            for _ in range(self.samples_per_category):
                data.append(self._generate_document(category))
                target.append(category)

        return data, np.array(target)


def create_text_datasets(num_categories=4):
    """创建训练集和测试集"""
    # 确保类别数在合理范围内
    num_categories = max(2, min(20, num_categories))

    # 创建训练集
    train_dataset = CustomTextDataset(
        num_categories=num_categories,
        samples_per_category=500,
        vocab_size=2000,
        random_state=42
    )

    # 创建测试集
    test_dataset = CustomTextDataset(
        num_categories=num_categories,
        samples_per_category=100,
        vocab_size=2000,
        random_state=43  # 不同的随机种子
    )

    return train_dataset, test_dataset


###########################################
# 第一部分:朴素贝叶斯文本分类
###########################################

class CustomMultinomialNB:
    """
    手写实现的多项式朴素贝叶斯分类器
    """

    def __init__(self, alpha=1.0):
        self.alpha = alpha  # 平滑参数
        self.class_priors = None  # 类别先验概率
        self.feature_probs = None  # 特征条件概率
        self.classes = None  # 类别列表
        self.n_features = None  # 特征数量
        self.feature_log_probs_ = None  # 特征对数概率(与sklearn兼容)
        self.class_log_prior_ = None  # 类别对数先验(与sklearn兼容)

    def fit(self, X, y):
        """
        训练朴素贝叶斯模型

        参数:
        X : scipy稀疏矩阵或numpy数组,形状为[n_samples, n_features]
            训练数据特征
        y : numpy数组,形状为[n_samples]
            训练数据标签

        返回:
        self : 对象实例
        """
        logger.info("开始训练自定义朴素贝叶斯模型...")
        start_time = time.time()

        # 获取唯一类别及其数量
        self.classes = np.unique(y)
        n_classes = len(self.classes)
        n_samples, self.n_features = X.shape

        # 计算类别先验概率
        class_counts = np.zeros(n_classes, dtype=np.float64)
        for i, c in enumerate(self.classes):
            class_counts[i] = np.sum(y == c)
        self.class_priors = class_counts / n_samples
        self.class_log_prior_ = np.log(self.class_priors)

        # 统计每个类别中各特征的出现次数
        feature_counts = np.zeros((n_classes, self.n_features), dtype=np.float64)

        # 高效处理稀疏矩阵
        if hasattr(X, "toarray"):  # 检查是否为稀疏矩阵
            # 使用tqdm显示进度
            for i in tqdm(range(n_samples), desc="处理样本", disable=not args.verbose):
                x_i = X[i].toarray().ravel()  # 获取第i个样本的特征向量
                c_idx = np.where(self.classes == y[i])[0][0]  # 获取该样本的类别索引
                feature_counts[c_idx] += x_i
        else:
            # 非稀疏矩阵的情况
            for c_idx, c in enumerate(self.classes):
                feature_counts[c_idx] = X[y == c].sum(axis=0)

        # 应用拉普拉斯平滑
        smoothed_fc = feature_counts + self.alpha

        # 计算每个特征在每个类别中的条件概率
        feature_totals = smoothed_fc.sum(axis=1).reshape(-1, 1)  # 每个类别的特征总数
        self.feature_probs = smoothed_fc / feature_totals
        self.feature_log_probs_ = np.log(self.feature_probs)

        logger.info(f"模型训练完成,用时 {time.time() - start_time:.2f} 秒")
        return self

    def predict_log_proba(self, X):
        """
        预测样本属于各类别的对数概率

        参数:
        X : scipy稀疏矩阵或numpy数组,形状为[n_samples, n_features]
            测试数据特征

        返回:
        numpy数组,形状为[n_samples, n_classes]
            每个样本属于各类别的对数概率
        """
        # 高效处理稀疏矩阵
        if hasattr(X, "toarray"):
            return self._predict_log_proba_sparse(X)

        # 计算每个类别的对数后验概率
        joint_log_likelihood = np.zeros((X.shape[0], len(self.classes)))
        for i, c in enumerate(self.classes):
            joint_log_likelihood[:, i] = self.class_log_prior_[i]
            joint_log_likelihood[:, i] += np.dot(X, self.feature_log_probs_[i])

        # 归一化后验概率
        log_prob_x = logsumexp(joint_log_likelihood, axis=1).reshape(-1, 1)
        log_probs = joint_log_likelihood - log_prob_x

        return log_probs

    def _predict_log_proba_sparse(self, X):
        """
        针对稀疏矩阵优化的对数概率计算
        """
        n_samples = X.shape[0]
        n_classes = len(self.classes)
        joint_log_likelihood = np.zeros((n_samples, n_classes))

        for i in range(n_samples):
            x_i = X[i].toarray().ravel()
            for j, c in enumerate(self.classes):
                joint_log_likelihood[i, j] = self.class_log_prior_[j]
                # 稀疏向量的乘法优化
                nonzero_indices = np.nonzero(x_i)[0]
                for idx in nonzero_indices:
                    joint_log_likelihood[i, j] += x_i[idx] * self.feature_log_probs_[j, idx]

        # 归一化后验概率
        log_prob_x = logsumexp(joint_log_likelihood, axis=1).reshape(-1, 1)
        log_probs = joint_log_likelihood - log_prob_x

        return log_probs

    def predict_proba(self, X):
        """
        预测样本属于各类别的概率

        参数:
        X : scipy稀疏矩阵或numpy数组,形状为[n_samples, n_features]
            测试数据特征

        返回:
        numpy数组,形状为[n_samples, n_classes]
            每个样本属于各类别的概率
        """
        return np.exp(self.predict_log_proba(X))

    def predict(self, X):
        """
        预测样本的类别

        参数:
        X : scipy稀疏矩阵或numpy数组,形状为[n_samples, n_features]
            测试数据特征

        返回:
        numpy数组,形状为[n_samples]
            预测的类别标签
        """
        # 返回概率最大的类别
        proba = self.predict_proba(X)
        return self.classes[np.argmax(proba, axis=1)]

    def score(self, X, y):
        """
        返回模型在给定测试数据和标签上的准确率

        参数:
        X : scipy稀疏矩阵或numpy数组,形状为[n_samples, n_features]
            测试数据特征
        y : numpy数组,形状为[n_samples]
            测试数据标签

        返回:
        float : 分类准确率
        """
        return accuracy_score(y, self.predict(X))


def logsumexp(arr, axis=0):
    """
    稳定计算log(sum(exp(x)))以避免数值溢出

    参数:
    arr : numpy数组
        输入数组
    axis : int
        沿着哪个轴计算和

    返回:
    numpy数组 : log(sum(exp(x)))的结果
    """
    max_val = np.max(arr, axis=axis, keepdims=True)
    ds = arr - max_val
    sum_exp = np.sum(np.exp(ds), axis=axis, keepdims=True)
    return max_val + np.log(sum_exp)


class EnsembleClassifier:
    """
    集成分类器,结合朴素贝叶斯、SVM和随机森林
    """

    def __init__(self, alpha=1.0):
        self.alpha = alpha
        self.custom_nb = CustomMultinomialNB(alpha=alpha)
        self.sklearn_nb = MultinomialNB(alpha=alpha)
        self.svm = SVC(probability=True, kernel='linear', C=1.0)
        self.rf = RandomForestClassifier(n_estimators=100, random_state=42)
        self.is_fitted = False

    def fit(self, X, y):
        """训练集成模型"""
        logger.info("训练集成分类器...")

        # 训练各个基分类器
        self.custom_nb.fit(X, y)
        self.sklearn_nb.fit(X, y)
        self.svm.fit(X, y)
        self.rf.fit(X, y)

        self.classes = np.unique(y)
        self.is_fitted = True

        return self

    def predict(self, X):
        """使用投票方式进行预测"""
        if not self.is_fitted:
            raise ValueError("模型尚未训练")

        # 获取各个模型的预测结果
        custom_nb_pred = self.custom_nb.predict(X)
        sklearn_nb_pred = self.sklearn_nb.predict(X)
        svm_pred = self.svm.predict(X)
        rf_pred = self.rf.predict(X)

        # 简单投票法(Majority Voting)
        predictions = np.vstack([custom_nb_pred, sklearn_nb_pred, svm_pred, rf_pred])
        majority_votes = np.zeros(X.shape[0], dtype=int)

        # 对每个样本统计投票结果
        for i in range(X.shape[0]):
            vote_count = {}
            for j in range(predictions.shape[0]):  # 遍历各模型
                vote = predictions[j, i]
                if vote in vote_count:
                    vote_count[vote] += 1
                else:
                    vote_count[vote] = 1

            # 找到票数最多的类别
            max_vote = 0
            max_class = None
            for cls, count in vote_count.items():
                if count > max_vote:
                    max_vote = count
                    max_class = cls

            majority_votes[i] = max_class

        return majority_votes

    def score(self, X, y):
        """计算集成模型的准确率"""
        return accuracy_score(y, self.predict(X))

    def get_feature_importances(self, X):
        """获取特征重要性"""
        # 使用随机森林的特征重要性
        return self.rf.feature_importances_


def run_text_classification(status_callback=None):
    """
    运行朴素贝叶斯文本分类实验

    参数:
    status_callback : 可选的回调函数,用于在GUI模式下更新状态

    返回:
    元组 : 包含分类结果的字典和可视化图表列表
    """
    if status_callback:
        status_callback("开始朴素贝叶斯文本分类实验...")

    logger.info("=" * 50)
    logger.info("开始朴素贝叶斯文本分类实验")
    logger.info("=" * 50)

    # 创建输出子目录
    nb_output_dir = os.path.join(args.output_dir, "naive_bayes")
    if args.save_results and not os.path.exists(nb_output_dir):
        os.makedirs(nb_output_dir)

    # 调整类别数量
    if args.categories > 20:
        args.categories = 20
    elif args.categories < 2:
        args.categories = 2

    num_categories = args.categories

    if args.interactive and not args.gui:
        try:
            num_categories = int(input("\n请输入分类类别数量 (2-20): "))
            num_categories = max(2, min(20, num_categories))
        except ValueError:
            print("使用默认类别数:", num_categories)

    # 加载自定义生成的数据集
    if status_callback:
        status_callback(f"生成模拟文本数据集,类别数量:{num_categories}...")

    logger.info(f"生成模拟文本数据集,类别数量:{num_categories}...")
    train_data, test_data = create_text_datasets(num_categories=num_categories)

    selected_categories = train_data.target_names
    logger.info(f"选择的类别: {selected_categories}")

    if status_callback:
        status_callback(f"训练集大小: {len(train_data.data)}, 测试集大小: {len(test_data.data)}")
    logger.info(f"训练集大小: {len(train_data.data)}, 测试集大小: {len(test_data.data)}")

    # 特征提取
    if status_callback:
        status_callback("特征提取中,使用TF-IDF向量化...")

    logger.info("特征提取中,使用TF-IDF向量化...")
    vectorizer = TfidfVectorizer(
        max_features=1000,
        min_df=2,
        ngram_range=(1, 2)  # 使用1-gram和2-gram提高特征表示能力
    )
    X_train = vectorizer.fit_transform(train_data.data)
    X_test = vectorizer.transform(test_data.data)
    y_train = train_data.target
    y_test = test_data.target

    logger.info(f"特征提取完成,特征维度: {X_train.shape[1]}")

    # 选择alpha参数
    alpha = args.alpha
    if args.interactive and not args.gui:
        try:
            alpha = float(input("\n请输入平滑参数alpha (推荐范围: 0.1 - 5.0): "))
            if alpha <= 0:
                print("alpha必须为正数,使用默认值1.0")
                alpha = 1.0
        except ValueError:
            print("输入无效,使用默认alpha=1.0")
            alpha = 1.0

    logger.info(f"使用平滑参数alpha={alpha}")

    # 训练并评估模型
    models = {}
    predictions = {}
    scores = {}
    visualizations = []

    # 训练自定义朴素贝叶斯模型
    if status_callback:
        status_callback("训练自定义朴素贝叶斯模型...")

    logger.info("训练自定义朴素贝叶斯模型...")
    custom_nb = CustomMultinomialNB(alpha=alpha)
    custom_nb.fit(X_train, y_train)
    custom_pred = custom_nb.predict(X_test)
    custom_accuracy = accuracy_score(y_test, custom_pred)
    custom_precision, custom_recall, custom_f1, _ = precision_recall_fscore_support(
        y_test, custom_pred, average='weighted'
    )

    models['custom_nb'] = custom_nb
    predictions['custom_nb'] = custom_pred
    scores['custom_nb'] = {
        'accuracy': custom_accuracy,
        'precision': custom_precision,
        'recall': custom_recall,
        'f1': custom_f1
    }

    logger.info(f"自定义模型评估结果:")
    logger.info(f"  准确率: {custom_accuracy:.4f}")
    logger.info(f"  精确率: {custom_precision:.4f}")
    logger.info(f"  召回率: {custom_recall:.4f}")
    logger.info(f"  F1分数: {custom_f1:.4f}")

    # 训练sklearn朴素贝叶斯模型
    if status_callback:
        status_callback("训练sklearn朴素贝叶斯模型...")

    logger.info("训练sklearn朴素贝叶斯模型...")
    sklearn_nb = MultinomialNB(alpha=alpha)
    sklearn_nb.fit(X_train, y_train)
    sklearn_pred = sklearn_nb.predict(X_test)
    sklearn_accuracy = accuracy_score(y_test, sklearn_pred)
    sklearn_precision, sklearn_recall, sklearn_f1, _ = precision_recall_fscore_support(
        y_test, sklearn_pred, average='weighted'
    )

    models['sklearn_nb'] = sklearn_nb
    predictions['sklearn_nb'] = sklearn_pred
    scores['sklearn_nb'] = {
        'accuracy': sklearn_accuracy,
        'precision': sklearn_precision,
        'recall': sklearn_recall,
        'f1': sklearn_f1
    }

    logger.info(f"sklearn模型评估结果:")
    logger.info(f"  准确率: {sklearn_accuracy:.4f}")
    logger.info(f"  精确率: {sklearn_precision:.4f}")
    logger.info(f"  召回率: {sklearn_recall:.4f}")
    logger.info(f"  F1分数: {sklearn_f1:.4f}")

    # 如果启用了集成学习,训练集成模型
    if args.ensemble or args.gui:
        if status_callback:
            status_callback("训练集成分类器...")

        logger.info("训练集成分类器...")
        ensemble_model = EnsembleClassifier(alpha=alpha)
        ensemble_model.fit(X_train, y_train)
        ensemble_pred = ensemble_model.predict(X_test)
        ensemble_accuracy = accuracy_score(y_test, ensemble_pred)
        ensemble_precision, ensemble_recall, ensemble_f1, _ = precision_recall_fscore_support(
            y_test, ensemble_pred, average='weighted'
        )

        models['ensemble'] = ensemble_model
        predictions['ensemble'] = ensemble_pred
        scores['ensemble'] = {
            'accuracy': ensemble_accuracy,
            'precision': ensemble_precision,
            'recall': ensemble_recall,
            'f1': ensemble_f1
        }

        logger.info(f"集成模型评估结果:")
        logger.info(f"  准确率: {ensemble_accuracy:.4f}")
        logger.info(f"  精确率: {ensemble_precision:.4f}")
        logger.info(f"  召回率: {ensemble_recall:.4f}")
        logger.info(f"  F1分数: {ensemble_f1:.4f}")

    # 参数网格搜索
    if args.verbose and (not args.interactive or args.gui):
        if status_callback:
            status_callback("进行参数网格搜索...")

        logger.info("进行参数网格搜索...")
        param_grid = {'alpha': [0.01, 0.1, 0.5, 1.0, 2.0, 5.0, 10.0]}
        grid_search = GridSearchCV(MultinomialNB(), param_grid, cv=5, scoring='accuracy')
        grid_search.fit(X_train, y_train)

        best_alpha = grid_search.best_params_['alpha']
        best_score = grid_search.best_score_
        logger.info(f"最佳alpha参数: {best_alpha}, 交叉验证准确率: {best_score:.4f}")

        # 使用最佳参数再次训练模型
        if best_alpha != alpha:
            logger.info(f"使用最佳alpha参数({best_alpha})训练新模型")
            best_nb = MultinomialNB(alpha=best_alpha)
            best_nb.fit(X_train, y_train)
            best_pred = best_nb.predict(X_test)
            best_accuracy = accuracy_score(y_test, best_pred)

            best_precision, best_recall, best_f1, _ = precision_recall_fscore_support(
                y_test, best_pred, average='weighted'
            )

            models['best_nb'] = best_nb
            predictions['best_nb'] = best_pred
            scores['best_nb'] = {
                'accuracy': best_accuracy,
                'precision': best_precision,
                'recall': best_recall,
                'f1': best_f1,
                'best_alpha': best_alpha
            }

            logger.info(f"最佳参数模型评估结果:")
            logger.info(f"  准确率: {best_accuracy:.4f}")
            logger.info(f"  精确率: {best_precision:.4f}")
            logger.info(f"  召回率: {best_recall:.4f}")
            logger.info(f"  F1分数: {best_f1:.4f}")

    # 创建可视化图表
    if status_callback:
        status_callback("生成可视化结果...")

    # 1. 混淆矩阵
    plt.figure(figsize=(10, 8))
    cm = confusion_matrix(y_test, custom_pred)
    ax = sns.heatmap(cm, annot=True, fmt="d", cmap="Blues",
                     xticklabels=selected_categories, yticklabels=selected_categories)
    plt.title("自定义模型混淆矩阵")
    plt.xlabel("预测标签")
    plt.ylabel("真实标签")
    plt.xticks(rotation=45, ha='right')
    plt.yticks(rotation=0)
    plt.tight_layout()

    # 保存图表
    if args.save_results:
        plt.savefig(os.path.join(nb_output_dir, 'confusion_matrix.png'), dpi=300)

    cm_fig = plt.gcf()
    visualizations.append(('confusion_matrix', cm_fig))

    # 2. 准确率比较
    plt.figure(figsize=(10, 6))
    model_names = list(scores.keys())
    accuracies = [scores[name]['accuracy'] for name in model_names]

    # 自定义模型名称标签
    model_labels = {
        'custom_nb': '自定义朴素贝叶斯',
        'sklearn_nb': 'Sklearn朴素贝叶斯',
        'ensemble': '集成模型',
        'best_nb': f'最优参数模型(α={scores.get("best_nb", {"best_alpha": alpha})["best_alpha"] if "best_nb" in scores else alpha})'
    }

    display_names = [model_labels.get(name, name) for name in model_names]

    plt.bar(display_names, accuracies, color=plt.cm.Paired(np.linspace(0, 1, len(model_names))))
    plt.title("模型准确率比较")
    plt.ylim([0, 1])
    for i, v in enumerate(accuracies):
        plt.text(i, v + 0.01, f"{v:.4f}", ha='center')
    plt.xticks(rotation=15, ha='right')
    plt.tight_layout()

    # 保存图表
    if args.save_results:
        plt.savefig(os.path.join(nb_output_dir, 'accuracy_comparison.png'), dpi=300)

    acc_fig = plt.gcf()
    visualizations.append(('accuracy_comparison', acc_fig))

    # 3. 性能指标对比
    plt.figure(figsize=(12, 8))
    metrics_names = ['准确率', '精确率', '召回率', 'F1分数']
    metrics_keys = ['accuracy', 'precision', 'recall', 'f1']

    x = np.arange(len(metrics_names))
    width = 0.8 / len(model_names)

    for i, model_name in enumerate(model_names):
        metric_values = [scores[model_name][key] for key in metrics_keys]
        plt.bar(x + (i - len(model_names) / 2 + 0.5) * width, metric_values, width,
                label=model_labels.get(model_name, model_name))

    plt.title("性能指标比较")
    plt.xticks(x, metrics_names)
    plt.ylim([0, 1])
    plt.legend()
    plt.tight_layout()

    # 保存图表
    if args.save_results:
        plt.savefig(os.path.join(nb_output_dir, 'metrics_comparison.png'), dpi=300)

    metrics_fig = plt.gcf()
    visualizations.append(('metrics_comparison', metrics_fig))

    # 4. 类别分布
    plt.figure(figsize=(12, 6))
    class_counts = np.bincount(y_train)
    plt.bar(selected_categories, class_counts, color=plt.cm.tab10.colors[:len(selected_categories)])
    plt.title("训练集类别分布")
    plt.xticks(rotation=45, ha='right')
    plt.tight_layout()

    # 保存图表
    if args.save_results:
        plt.savefig(os.path.join(nb_output_dir, 'class_distribution.png'), dpi=300)

    class_fig = plt.gcf()
    visualizations.append(('class_distribution', class_fig))

    # 保存结果
    if args.save_results:
        # 保存详细分类报告
        with open(os.path.join(nb_output_dir, 'classification_report.txt'), 'w', encoding='utf-8') as f:
            f.write("朴素贝叶斯文本分类实验结果\n")
            f.write("=" * 50 + "\n\n")
            f.write("实验参数:\n")
            f.write(f"- 类别数量: {len(selected_categories)}\n")
            f.write(f"- 平滑参数alpha: {alpha}\n")
            f.write(f"- 训练集大小: {len(train_data.data)}\n")
            f.write(f"- 测试集大小: {len(test_data.data)}\n")
            f.write(f"- 特征维度: {X_train.shape[1]}\n\n")

            for model_name, pred in predictions.items():
                f.write(f"{model_labels.get(model_name, model_name)}分类报告:\n")
                f.write(classification_report(y_test, pred, target_names=selected_categories))
                f.write("\n\n")

            if "best_nb" in scores:
                f.write("\n参数网格搜索结果:\n")
                f.write(f"最佳alpha参数: {scores['best_nb']['best_alpha']}\n")

        # 保存模型
        for name, model in models.items():
            with open(os.path.join(nb_output_dir, f"{name}_model.pkl"), 'wb') as f:
                pickle.dump(model, f)

        # 保存向量化器
        with open(os.path.join(nb_output_dir, "vectorizer.pkl"), 'wb') as f:
            pickle.dump(vectorizer, f)

    if not args.interactive and not args.gui:
        plt.show()
    elif args.interactive and not args.gui:
        print("\n模型比较结果:")
        for name in model_names:
            print(f"{model_labels.get(name, name)}:")
            print(f"   - 准确率: {scores[name]['accuracy']:.4f}")
            print(f"   - 精确率: {scores[name]['precision']:.4f}")
            print(f"   - 召回率: {scores[name]['recall']:.4f}")
            print(f"   - F1分数: {scores[name]['f1']:.4f}")

        # 展示混淆矩阵
        show_option = input("\n是否显示可视化结果? (y/n): ")
        if show_option.lower() == 'y':
            plt.show()

    if status_callback:
        status_callback("朴素贝叶斯文本分类实验完成")

    logger.info("朴素贝叶斯文本分类实验完成")

    # 返回模型性能和可视化图表
    results = {
        'models': models,
        'predictions': predictions,
        'scores': scores,
        'selected_categories': selected_categories,
        'X_train': X_train,
        'X_test': X_test,
        'y_train': y_train,
        'y_test': y_test,
        'vectorizer': vectorizer
    }

    return results, visualizations


###########################################
# 第二部分:马尔可夫网络图像去噪
###########################################

class MRFImageDenoiser:
    """
    使用马尔可夫随机场(MRF)进行图像去噪的类
    """

    def __init__(self, eta=2.0, beta=1.0, max_iters=10):
        """
        初始化MRF去噪器

        参数:
        eta : float, 默认=2.0
            数据项权重(噪声图像和去噪图像的关联强度)
        beta : float, 默认=1.0
            光滑项权重(相邻像素相似性的影响强度)
        max_iters : int, 默认=10
            ICM算法的最大迭代次数
        """
        self.eta = eta
        self.beta = beta
        self.max_iters = max_iters
        self.noisy_img = None
        self.denoised_img = None
        self.energies = []  # 记录每次迭代的能量
        self.intermediate_results = []  # 记录中间结果图像

    def compute_energy(self, x, y, i, j):
        """
        计算像素(i,j)处的局部能量

        参数:
        x : 2D numpy数组
            当前图像
        y : 2D numpy数组
            噪声图像
        i, j : int
            像素坐标

        返回:
        float : 该像素处的局部能量
        """
        # 数据项
        energy = -self.eta * x[i, j] * y[i, j]

        # 光滑项(与相邻像素的相互作用)
        for di, dj in [(-1, 0), (1, 0), (0, -1), (0, 1)]:
            ni, nj = i + di, j + dj
            if 0 <= ni < x.shape[0] and 0 <= nj < x.shape[1]:
                energy -= self.beta * x[i, j] * x[ni, nj]

        return energy

    def compute_total_energy(self, x, y):
        """
        计算整个图像的总能量

        参数:
        x : 2D numpy数组
            当前图像
        y : 2D numpy数组
            噪声图像

        返回:
        float : 总能量
        """
        total_energy = 0
        height, width = x.shape

        # 数据项
        total_energy -= self.eta * np.sum(x * y)

        # 光滑项
        # 水平相邻像素
        total_energy -= self.beta * np.sum(x[:, :-1] * x[:, 1:])
        # 垂直相邻像素
        total_energy -= self.beta * np.sum(x[:-1, :] * x[1:, :])

        return total_energy

    def icm_update(self, x, y, callback=None):
        """
        使用ICM(迭代条件模式)算法更新图像

        参数:
        x : 2D numpy数组
            当前图像
        y : 2D numpy数组
            噪声图像
        callback : 回调函数
            用于GUI更新

        返回:
        2D numpy数组 : 更新后的图像
        bool : 是否有像素被更新
        """
        height, width = x.shape
        updated = False

        if args.parallel and height * width > 10000:
            # 并行处理大图像(将图像分成块并行处理)
            return self._parallel_icm_update(x, y, callback)

        # 使用tqdm显示进度
        progress_update = args.verbose and height * width > 10000

        # 创建像素坐标列表
        pixels = [(i, j) for i in range(height) for j in range(width)]

        # 使用tqdm进度条(仅在命令行模式下)
        if not args.gui:
            iter_pixels = tqdm(pixels, desc="ICM更新", disable=not progress_update)
        else:
            iter_pixels = pixels

            # 用于GUI进度更新
            total_pixels = len(pixels)
            processed_pixels = 0

        for i, j in iter_pixels:
            # 尝试将当前像素取+1,计算能量
            x[i, j] = 1
            E_plus = self.compute_energy(x, y, i, j)

            # 尝试将当前像素取-1,计算能量
            x[i, j] = -1
            E_minus = self.compute_energy(x, y, i, j)

            # 选择能量较小的状态
            if E_plus < E_minus:
                x[i, j] = 1
                if E_plus != E_minus:
                    updated = True
            else:
                x[i, j] = -1
                if E_plus != E_minus:
                    updated = True

            # GUI模式下更新进度
            if args.gui and callback:
                processed_pixels += 1
                if processed_pixels % 1000 == 0 or processed_pixels == total_pixels:
                    progress = processed_pixels / total_pixels * 100
                    callback(f"ICM更新进度: {progress:.1f}%")

        return x, updated

    def _parallel_icm_update(self, x, y, callback=None, block_size=100):
        """
        并行版本的ICM更新

        参数:
        x : 2D numpy数组
            当前图像
        y : 2D numpy数组
            噪声图像
        callback : 回调函数
            用于GUI更新
        block_size : int
            处理块大小

        返回:
        2D numpy数组 : 更新后的图像
        bool : 是否有像素被更新
        """
        height, width = x.shape
        result_x = np.copy(x)
        updated = False

        # 定义需要处理的块
        blocks = []
        for i_start in range(0, height, block_size):
            for j_start in range(0, width, block_size):
                i_end = min(i_start + block_size, height)
                j_end = min(j_start + block_size, width)
                blocks.append((i_start, i_end, j_start, j_end))

        # 定义处理单个块的函数
        def process_block(block):
            i_start, i_end, j_start, j_end = block
            block_x = np.copy(x[i_start:i_end, j_start:j_end])
            block_y = y[i_start:i_end, j_start:j_end]
            block_updated = False
            local_result_x = np.copy(result_x)

            for i_rel in range(i_end - i_start):
                for j_rel in range(j_end - j_start):
                    i, j = i_start + i_rel, j_start + j_rel

                    # 计算当前像素在x中的全局坐标
                    i_glob, j_glob = i, j

                    # 尝试将当前像素取+1,计算能量
                    local_result_x[i_glob, j_glob] = 1
                    E_plus = self.compute_energy(local_result_x, y, i_glob, j_glob)

                    # 尝试将当前像素取-1,计算能量
                    local_result_x[i_glob, j_glob] = -1
                    E_minus = self.compute_energy(local_result_x, y, i_glob, j_glob)

                    # 选择能量较小的状态
                    if E_plus < E_minus:
                        local_result_x[i_glob, j_glob] = 1
                        if E_plus != E_minus:
                            block_updated = True
                    else:
                        local_result_x[i_glob, j_glob] = -1
                        if E_plus != E_minus:
                            block_updated = True

            return (block, local_result_x[i_start:i_end, j_start:j_end], block_updated)

        # 并行处理块,如果multiprocessing可用
        try:
            # 计算并行核心数
            n_cores = min(cpu_count(), 4) if cpu_count() > 1 else 1

            with Pool(processes=n_cores) as pool:
                # 使用tqdm进度条(仅在命令行模式下)
                if not args.gui:
                    results = list(tqdm(
                        pool.imap(process_block, blocks),
                        total=len(blocks),
                        desc="并行ICM更新",
                        disable=not args.verbose
                    ))
                else:
                    # 用于GUI进度更新
                    total_blocks = len(blocks)
                    processed_blocks = 0

                    # 创建回调函数更新进度
                    def callback_wrapper(result):
                        nonlocal processed_blocks
                        processed_blocks += 1
                        if callback and (processed_blocks % 10 == 0 or processed_blocks == total_blocks):
                            progress = processed_blocks / total_blocks * 100
                            callback(f"并行ICM更新: {progress:.1f}%")

                    # 初始化异步结果列表
                    async_results = []
                    for block in blocks:
                        res = pool.apply_async(process_block, (block,), callback=callback_wrapper)
                        async_results.append(res)

                    # 等待所有结果完成
                    results = [r.get() for r in async_results]

            # 更新图像并检查是否有更新
            for (i_start, i_end, j_start, j_end), block_result, block_updated in results:
                result_x[i_start:i_end, j_start:j_end] = block_result
                if block_updated:
                    updated = True

        except (ImportError, ValueError, RuntimeError) as e:
            # 如果出现并行处理问题,退回到串行处理
            logger.warning(f"并行处理失败: {str(e)},使用串行处理")

            # 使用tqdm进度条(仅在命令行模式下)
            if not args.gui:
                block_iterator = tqdm(blocks, desc="串行ICM更新", disable=not args.verbose)
            else:
                block_iterator = blocks
                total_blocks = len(blocks)
                processed_blocks = 0

            for block in block_iterator:
                i_start, i_end, j_start, j_end = block
                _, block_result, block_updated = process_block(block)
                result_x[i_start:i_end, j_start:j_end] = block_result
                if block_updated:
                    updated = True

                # GUI模式下更新进度
                if args.gui and callback:
                    processed_blocks += 1
                    if processed_blocks % 10 == 0 or processed_blocks == total_blocks:
                        progress = processed_blocks / total_blocks * 100
                        callback(f"串行块处理: {progress:.1f}%")

        return result_x, updated

    def denoise(self, noisy_img, callback=None):
        """
        对噪声图像进行去噪

        参数:
        noisy_img : 2D numpy数组
            噪声图像,值为-1(黑)或1(白)
        callback : 回调函数
            用于GUI更新进度和中间结果

        返回:
        2D numpy数组 : 去噪后的图像
        """
        if callback:
            callback(f"开始去噪,参数:eta={self.eta}, beta={self.beta}, max_iters={self.max_iters}")

        logger.info(f"开始去噪,参数:eta={self.eta}, beta={self.beta}, max_iters={self.max_iters}")
        self.noisy_img = np.copy(noisy_img)
        self.denoised_img = np.copy(noisy_img)  # 初始化为噪声图像

        # 记录初始能量和初始图像
        init_energy = self.compute_total_energy(self.denoised_img, self.noisy_img)
        self.energies = [init_energy]
        self.intermediate_results = [np.copy(noisy_img)]

        # ICM迭代优化
        for iter_num in range(self.max_iters):
            if callback:
                callback(f"进行第 {iter_num + 1}/{self.max_iters} 次迭代...")

            start_time = time.time()

            # ICM更新
            self.denoised_img, updated = self.icm_update(self.denoised_img, self.noisy_img, callback)

            # 保存当前迭代结果
            self.intermediate_results.append(np.copy(self.denoised_img))

            # 计算当前能量
            current_energy = self.compute_total_energy(self.denoised_img, self.noisy_img)
            self.energies.append(current_energy)

            iter_time = time.time() - start_time
            logger.info(f"迭代 {iter_num + 1}/{self.max_iters}, 能量: {current_energy:.2f}, 用时: {iter_time:.2f}秒")

            if callback:
                callback(f"迭代 {iter_num + 1} 完成,能量: {current_energy:.2f},用时: {iter_time:.2f}秒")

            # 如果没有像素更新,提前终止
            if not updated:
                logger.info(f"能量稳定,提前终止在第{iter_num + 1}次迭代")
                if callback:
                    callback(f"能量稳定,提前终止在第{iter_num + 1}次迭代")
                break

        final_energy = self.energies[-1]
        energy_reduction = (init_energy - final_energy) / abs(init_energy) * 100 if init_energy != 0 else 0
        logger.info(f"去噪完成,能量从{init_energy:.2f}降至{final_energy:.2f},减少了{energy_reduction:.2f}%")

        if callback:
            callback(f"去噪完成,能量从{init_energy:.2f}降至{final_energy:.2f},减少了{energy_reduction:.2f}%")

        return self.denoised_img

    def compute_metrics(self, original_img, noisy_img, denoised_img):
        """
        计算去噪效果的评估指标

        参数:
        original_img : 2D numpy数组
            原始干净图像
        noisy_img : 2D numpy数组
            噪声图像
        denoised_img : 2D numpy数组
            去噪后的图像

        返回:
        dict : 包含各种评估指标的字典
        """
        # 计算错误率(与原始图像不一致的像素比例)
        original_error = np.sum(noisy_img != original_img) / original_img.size
        denoised_error = np.sum(denoised_img != original_img) / original_img.size
        error_reduction = (original_error - denoised_error) / original_error * 100 if original_error != 0 else 0

        # 计算平均绝对差
        original_mad = np.mean(np.abs(noisy_img - original_img))
        denoised_mad = np.mean(np.abs(denoised_img - original_img))

        # 计算PSNR(以-1/1图像为基础,转换到0-1范围)
        def compute_psnr(img1, img2):
            img1_01 = (img1 + 1) / 2  # 转换到0-1范围
            img2_01 = (img2 + 1) / 2
            mse = np.mean((img1_01 - img2_01) ** 2)
            if mse == 0:
                return float('inf')
            return 10 * np.log10(1.0 / mse)

        noisy_psnr = compute_psnr(original_img, noisy_img)
        denoised_psnr = compute_psnr(original_img, denoised_img)
        psnr_improvement = denoised_psnr - noisy_psnr

        metrics = {
            "原始错误率": original_error,
            "去噪后错误率": denoised_error,
            "错误率减少百分比": error_reduction,
            "原始MAD": original_mad,
            "去噪后MAD": denoised_mad,
            "原始PSNR": noisy_psnr,
            "去噪后PSNR": denoised_psnr,
            "PSNR提升": psnr_improvement
        }

        return metrics


def create_synthetic_image(height=100, width=100, pattern_type="checkerboard"):
    """
    创建合成二值测试图像

    参数:
    height, width : int
        图像尺寸
    pattern_type : str
        图像模式类型,可选值:
        - "checkerboard": 棋盘格
        - "horizontal_stripes": 水平条纹
        - "vertical_stripes": 垂直条纹
        - "circle": 中心圆形
        - "cross": 十字形
        - "random": 随机生成
        - "text": MRF文本图案

    返回:
    2D numpy数组 : 合成图像,值为-1(黑)或1(白)
    """
    img = np.ones((height, width), dtype=int)

    if pattern_type == "checkerboard":
        checker_size = min(height, width) // 10
        for i in range(height):
            for j in range(width):
                if ((i // checker_size) + (j // checker_size)) % 2 == 0:
                    img[i, j] = -1

    elif pattern_type == "horizontal_stripes":
        stripe_width = height // 10
        for i in range(height):
            if (i // stripe_width) % 2 == 0:
                img[i, :] = -1

    elif pattern_type == "vertical_stripes":
        stripe_width = width // 10
        for j in range(width):
            if (j // stripe_width) % 2 == 0:
                img[:, j] = -1

    elif pattern_type == "circle":
        center_y, center_x = height // 2, width // 2
        radius = min(height, width) // 3
        for i in range(height):
            for j in range(width):
                if (i - center_y) ** 2 + (j - center_x) ** 2 <= radius ** 2:
                    img[i, j] = -1

    elif pattern_type == "cross":
        cross_width = min(height, width) // 10
        center_y, center_x = height // 2, width // 2
        # 垂直线
        img[center_y - cross_width:center_y + cross_width, center_x - cross_width // 2:center_x + cross_width // 2] = -1
        # 水平线
        img[center_y - cross_width // 2:center_y + cross_width // 2, center_x - cross_width:center_x + cross_width] = -1

    elif pattern_type == "random":
        # 使用随机块创建更有意义的结构
        block_size = min(height, width) // 10
        for _ in range(5):  # 创建5个随机块
            block_y = np.random.randint(0, height - block_size)
            block_x = np.random.randint(0, width - block_size)
            img[block_y:block_y + block_size, block_x:block_x + block_size] = -1

    elif pattern_type == "text":
        # 创建文本"MRF"的简单点阵图,作为艺术字效果
        if height >= 40 and width >= 100:
            # "M"字形
            for i in range(10, 30):
                j_mid = width // 4
                j1 = max(j_mid - (i - 10), j_mid - 10)
                j2 = min(j_mid + (i - 10), j_mid + 10)
                img[i, j1:j2 + 1] = -1

            # "R"字形
            r_start = width // 2 - 5
            for i in range(10, 30):
                if i < 20:  # R的上半部分(半圆形)
                    r_width = 10 - abs(i - 15)
                    img[i, r_start:r_start + 5 + r_width] = -1
                else:  # R的下半部分(斜线)
                    r_width = i - 20
                    img[i, r_start:r_start + 5 + r_width] = -1

            # "F"字形
            f_start = 3 * width // 4
            for i in range(10, 30):
                if i == 10 or i == 20:  # F的横线部分
                    img[i, f_start:f_start + 15] = -1
                else:  # F的竖线部分
                    img[i, f_start:f_start + 5] = -1

    return img


def add_noise(img, noise_rate):
    """
    向图像添加噪声

    参数:
    img : 2D numpy数组
        原始图像
    noise_rate : float
        噪声比例(0.0-1.0)

    返回:
    2D numpy数组 : 添加噪声后的图像
    """
    noisy = np.copy(img)
    height, width = img.shape
    num_pixels = height * width
    num_noise = int(noise_rate * num_pixels)

    # 随机选择像素位置
    noise_indices = np.random.choice(num_pixels, num_noise, replace=False)
    noise_coords = np.unravel_index(noise_indices, (height, width))

    # 翻转选中像素的值(-1变1,1变-1)
    noisy[noise_coords] = -noisy[noise_coords]

    return noisy


def run_image_denoising(status_callback=None):
    """
    运行马尔可夫网络图像去噪实验

    参数:
    status_callback : 回调函数
        用于GUI更新状态

    返回:
    字典 : 包含去噪结果和指标的字典
    列表 : 可视化图表列表
    """
    if status_callback:
        status_callback("开始马尔可夫网络图像去噪实验...")

    logger.info("=" * 50)
    logger.info("开始马尔可夫网络图像去噪实验")
    logger.info("=" * 50)

    # 创建输出子目录
    mrf_output_dir = os.path.join(args.output_dir, "mrf_denoising")
    if args.save_results and not os.path.exists(mrf_output_dir):
        os.makedirs(mrf_output_dir)

    # 设置参数
    noise_rate = args.noise_rate
    eta = args.eta
    beta = args.beta
    max_iters = args.max_iters

    if args.interactive and not args.gui:
        # 交互式设置参数
        print("\n马尔可夫随机场图像去噪")

        # 选择图像类型
        pattern_types = ["checkerboard", "horizontal_stripes", "vertical_stripes",
                         "circle", "cross", "random", "text"]
        print("\n可用的图像模式:")
        for i, pat in enumerate(pattern_types):
            print(f"{i + 1}. {pat}")

        pattern_idx = 0
        while True:
            try:
                pattern_idx = int(input("\n请选择图像模式(1-7): ")) - 1
                if 0 <= pattern_idx < len(pattern_types):
                    break
                print("无效选择,请输入1-7之间的数字!")
            except ValueError:
                print("请输入有效的数字!")
                pattern_idx = 0  # 默认选择第一个
                break

        pattern_type = pattern_types[pattern_idx]

        # 设置图像尺寸
        size = 100
        try:
            size = int(input("\n请输入图像尺寸(建议范围: 50-200): "))
            if size < 10:
                size = 10
            elif size > 500:
                print("警告: 大尺寸图像可能处理较慢")
                size = min(size, 500)
        except ValueError:
            print("使用默认尺寸100")
            size = 100

        # 设置噪声率
        while True:
            try:
                noise_rate = float(input("\n请输入噪声率(0.0-0.5): "))
                if 0 <= noise_rate <= 0.5:
                    break
                print("噪声率应在0.0-0.5范围内!")
            except ValueError:
                print("请输入有效的小数!")
                noise_rate = 0.1  # 默认值
                break

        # 设置模型参数
        try:
            eta = float(input("\n请输入数据项权重eta(建议范围: 1.0-5.0): "))
            beta = float(input("\n请输入光滑项权重beta(建议范围: 0.5-2.0): "))
            max_iters = int(input("\n请输入最大迭代次数(建议范围: 5-20): "))
        except ValueError:
            print("使用默认参数设置")

        # 创建原始图像和噪声图像
        original_img = create_synthetic_image(size, size, pattern_type)
        noisy_img = add_noise(original_img, noise_rate)
    else:
        # 非交互模式或GUI模式,使用默认图像
        pattern_type = "text" if args.gui else "checkerboard"
        size = 100
        original_img = create_synthetic_image(size, size, pattern_type)
        noisy_img = add_noise(original_img, noise_rate)

    # 显示原始图像和噪声图像
    if status_callback:
        status_callback(f"创建了大小为{original_img.shape}的原始图像,添加了{noise_rate * 100:.1f}%的噪声")

    logger.info(f"创建了大小为{original_img.shape}的原始图像,添加了{noise_rate * 100:.1f}%的噪声")

    # 计算噪声率
    error_rate = np.sum(noisy_img != original_img) / original_img.size
    logger.info(f"实际噪声率: {error_rate * 100:.2f}%")

    if status_callback:
        status_callback(f"实际噪声率: {error_rate * 100:.2f}%")

    # 创建去噪器并进行去噪
    denoiser = MRFImageDenoiser(eta=eta, beta=beta, max_iters=max_iters)
    denoised_img = denoiser.denoise(noisy_img, status_callback)

    # 计算评估指标
    metrics = denoiser.compute_metrics(original_img, noisy_img, denoised_img)

    for name, value in metrics.items():
        if 'PSNR' in name:
            logger.info(f"{name}: {value:.2f} dB")
        elif '百分比' in name:
            logger.info(f"{name}: {value:.2f}%")
        else:
            logger.info(f"{name}: {value:.4f}")

    # 创建可视化结果
    visualizations = []

    if status_callback:
        status_callback("生成可视化结果...")

    # 1. 原始图像
    plt.figure(figsize=(6, 6))
    plt.imshow(original_img, cmap='binary_r')
    plt.title("原始图像")
    plt.axis('off')
    plt.tight_layout()

    # 保存图表
    if args.save_results:
        plt.savefig(os.path.join(mrf_output_dir, 'original_image.png'), dpi=300)

    orig_fig = plt.gcf()
    visualizations.append(('original_image', orig_fig))

    # 2. 噪声图像
    plt.figure(figsize=(6, 6))
    plt.imshow(noisy_img, cmap='binary_r')
    plt.title(f"噪声图像 (噪声率: {error_rate * 100:.2f}%)")
    plt.axis('off')
    plt.tight_layout()

    # 保存图表
    if args.save_results:
        plt.savefig(os.path.join(mrf_output_dir, 'noisy_image.png'), dpi=300)

    noisy_fig = plt.gcf()
    visualizations.append(('noisy_image', noisy_fig))

    # 3. 去噪后图像
    plt.figure(figsize=(6, 6))
    plt.imshow(denoised_img, cmap='binary_r')
    plt.title(f"去噪后图像 (错误率: {metrics['去噪后错误率'] * 100:.2f}%)")
    plt.axis('off')
    plt.tight_layout()

    # 保存图表
    if args.save_results:
        plt.savefig(os.path.join(mrf_output_dir, 'denoised_image.png'), dpi=300)

    denoised_fig = plt.gcf()
    visualizations.append(('denoised_image', denoised_fig))

    # 4. 错误对比图
    plt.figure(figsize=(6, 6))
    noisy_error = noisy_img != original_img
    denoised_error = denoised_img != original_img

    # 创建RGB差异图,红色表示噪声图像中的错误,蓝色表示去噪后的错误,紫色表示两者都有错误
    diff_img = np.zeros((*original_img.shape, 3))
    diff_img[noisy_error, 0] = 1.0  # 红色通道表示噪声图像的错误
    diff_img[denoised_error, 2] = 1.0  # 蓝色通道表示去噪后的错误

    plt.imshow(diff_img)
    plt.title("错误对比 (红:噪声错误, 蓝:去噪后错误, 紫:共同错误)")
    plt.axis('off')
    plt.tight_layout()

    # 保存图表
    if args.save_results:
        plt.savefig(os.path.join(mrf_output_dir, 'error_comparison.png'), dpi=300)

    diff_fig = plt.gcf()
    visualizations.append(('error_comparison', diff_fig))

    # 5. 能量变化曲线
    plt.figure(figsize=(10, 5))
    plt.plot(denoiser.energies, 'o-', linewidth=2)
    plt.title('能量变化曲线')
    plt.xlabel('迭代次数')
    plt.ylabel('能量')
    plt.grid(True)
    plt.tight_layout()

    # 保存图表
    if args.save_results:
        plt.savefig(os.path.join(mrf_output_dir, 'energy_curve.png'), dpi=300)

    energy_fig = plt.gcf()
    visualizations.append(('energy_curve', energy_fig))

    # 6. 迭代过程动画帧
    iter_figs = []
    if args.save_results:
        os.makedirs(os.path.join(mrf_output_dir, 'iterations'), exist_ok=True)

    for i, img in enumerate(denoiser.intermediate_results):
        plt.figure(figsize=(6, 6))
        plt.imshow(img, cmap='binary_r')
        if i == 0:
            plt.title(f"初始噪声图像")
        else:
            plt.title(f"迭代 {i}/{max_iters} 后的结果")
        plt.axis('off')
        plt.tight_layout()

        # 保存迭代图
        if args.save_results:
            plt.savefig(os.path.join(mrf_output_dir, f'iterations/iter_{i:02d}.png'), dpi=300)

        iter_fig = plt.gcf()
        iter_figs.append(iter_fig)
        visualizations.append((f'iteration_{i}', iter_fig))

    # 7. 参数影响图 (beta和eta的组合)
    if args.verbose or args.gui:
        # 测试不同的参数组合
        plt.figure(figsize=(12, 10))
        param_combinations = [
            (1.0, 1.0),  # (eta, beta)
            (2.0, 1.0),
            (3.0, 1.0),
            (2.0, 0.5),
            (2.0, 2.0)
        ]

        # 为每个参数组合创建独立的图表
        param_figs = []

        # 快速测试模式,只迭代3次
        test_iters = 3

        for idx, (test_eta, test_beta) in enumerate(param_combinations):
            if status_callback:
                status_callback(
                    f"测试参数组合 {idx + 1}/{len(param_combinations)}: eta={test_eta}, beta={test_beta}...")

            test_denoiser = MRFImageDenoiser(eta=test_eta, beta=test_beta, max_iters=test_iters)
            test_result = test_denoiser.denoise(noisy_img)
            test_metrics = test_denoiser.compute_metrics(original_img, noisy_img, test_result)

            # 创建单独的图表
            plt.figure(figsize=(6, 6))
            plt.imshow(test_result, cmap='binary_r')
            plt.title(f"eta={test_eta}, beta={test_beta}\nPSNR: {test_metrics['去噪后PSNR']:.2f} dB")
            plt.axis('off')
            plt.tight_layout()

            # 保存图表
            if args.save_results:
                plt.savefig(os.path.join(mrf_output_dir, f'param_eta{test_eta}_beta{test_beta}.png'), dpi=300)

            param_fig = plt.gcf()
            param_figs.append(param_fig)
            visualizations.append((f'param_eta{test_eta}_beta{test_beta}', param_fig))

    # 8. 参考图 - 所有参数组合比较
    if args.verbose or args.gui:
        plt.figure(figsize=(12, 8))
        plt.subplot(2, 3, 1)
        plt.imshow(original_img, cmap='binary_r')
        plt.title("原始图像(参考)")
        plt.axis('off')

        for idx, (test_eta, test_beta) in enumerate(param_combinations):
            plt.subplot(2, 3, idx + 2)
            plt.imshow(param_figs[idx].get_axes()[0].get_images()[0].get_array(), cmap='binary_r')
            plt.title(f"eta={test_eta}, beta={test_beta}")
            plt.axis('off')

        plt.tight_layout()

        # 保存图表
        if args.save_results:
            plt.savefig(os.path.join(mrf_output_dir, 'parameter_comparison.png'), dpi=300)

        params_comparison_fig = plt.gcf()
        visualizations.append(('parameter_comparison', params_comparison_fig))

    # 保存结果
    if args.save_results:
        # 保存原始、噪声和去噪后的图像数据
        np.save(os.path.join(mrf_output_dir, 'original_img.npy'), original_img)
        np.save(os.path.join(mrf_output_dir, 'noisy_img.npy'), noisy_img)
        np.save(os.path.join(mrf_output_dir, 'denoised_img.npy'), denoised_img)

        # 保存详细评估指标
        with open(os.path.join(mrf_output_dir, 'denoising_metrics.txt'), 'w', encoding='utf-8') as f:
            f.write("马尔可夫网络图像去噪实验结果\n")
            f.write("=" * 50 + "\n\n")
            f.write("实验参数:\n")
            f.write(f"- 图像尺寸: {original_img.shape}\n")
            f.write(f"- 图像模式: {pattern_type}\n")
            f.write(f"- 噪声率: {noise_rate}\n")
            f.write(f"- 数据项权重(eta): {eta}\n")
            f.write(f"- 光滑项权重(beta): {beta}\n")
            f.write(f"- 最大迭代次数: {max_iters}\n\n")

            f.write("评估指标:\n")
            for name, value in metrics.items():
                if 'PSNR' in name:
                    f.write(f"- {name}: {value:.2f} dB\n")
                elif '百分比' in name:
                    f.write(f"- {name}: {value:.2f}%\n")
                else:
                    f.write(f"- {name}: {value:.4f}\n")

    if not args.interactive and not args.gui:
        plt.show()
    elif args.interactive and not args.gui:
        print("\n去噪结果:")
        for name, value in metrics.items():
            if 'PSNR' in name:
                print(f"- {name}: {value:.2f} dB")
            elif '百分比' in name:
                print(f"- {name}: {value:.2f}%")
            else:
                print(f"- {name}: {value:.4f}")

        show_option = input("\n是否显示可视化结果? (y/n): ")
        if show_option.lower() == 'y':
            plt.show()

    if status_callback:
        status_callback("马尔可夫网络图像去噪实验完成")

    logger.info("马尔可夫网络图像去噪实验完成")

    # 返回结果
    results = {
        'original_img': original_img,
        'noisy_img': noisy_img,
        'denoised_img': denoised_img,
        'metrics': metrics,
        'energies': denoiser.energies,
        'intermediate_results': denoiser.intermediate_results,
        'params': {
            'eta': eta,
            'beta': beta,
            'max_iters': max_iters,
            'noise_rate': noise_rate,
            'pattern_type': pattern_type
        }
    }

    return results, visualizations


###########################################
# GUI部分
###########################################

class TextRedirector:
    """
    用于重定向stdout到Tkinter文本控件
    """

    def __init__(self, text_widget):
        self.text_widget = text_widget
        self.buffer = ""

    def write(self, text):
        self.buffer += text
        self.text_widget.configure(state="normal")
        self.text_widget.insert(tk.END, text)
        self.text_widget.see(tk.END)
        self.text_widget.configure(state="disabled")

    def flush(self):
        pass


class FigureWindow(tk.Toplevel):
    """
    独立的图形窗口,用于显示图表
    """

    def __init__(self, parent, fig, title="图形窗口"):
        super().__init__(parent)
        self.title(title)
        self.geometry("800x600")

        # 创建画布
        self.canvas = FigureCanvasTkAgg(fig, self)
        self.canvas.draw()

        # 添加导航工具栏
        toolbar_frame = ttk.Frame(self)
        toolbar_frame.pack(side=tk.TOP, fill=tk.X)
        toolbar = NavigationToolbar2Tk(self.canvas, toolbar_frame)
        toolbar.update()

        # 将画布放入窗口
        self.canvas.get_tk_widget().pack(fill=tk.BOTH, expand=True)

        # 添加保存按钮
        btn_frame = ttk.Frame(self)
        btn_frame.pack(side=tk.BOTTOM, fill=tk.X, pady=5)

        save_btn = ttk.Button(btn_frame, text="保存图像", command=self.save_figure)
        save_btn.pack(side=tk.RIGHT, padx=10)

    def save_figure(self):
        """保存图像到文件"""
        file_path = filedialog.asksaveasfilename(
            defaultextension=".png",
            filetypes=[("PNG图像", "*.png"), ("JPEG图像", "*.jpg"), ("SVG矢量图", "*.svg"), ("PDF文档", "*.pdf")]
        )
        if file_path:
            self.canvas.figure.savefig(file_path, dpi=300)
            messagebox.showinfo("保存成功", f"图像已保存到:\n{file_path}")


class App(tk.Tk):
    """
    主应用程序GUI
    """

    def __init__(self):
        super().__init__()

        self.title("Python机器学习:朴素贝叶斯文本分类与马尔可夫网络图像去噪实现")
        self.geometry("1024x768")

        # 设置标题和说明
        title_frame = ttk.Frame(self)
        title_frame.pack(fill=tk.X, padx=10, pady=10)

        title_label = ttk.Label(
            title_frame,
            text="Python机器学习工具箱",
            font=("Arial", 16, "bold")
        )
        title_label.pack()

        subtitle_label = ttk.Label(
            title_frame,
            text="朴素贝叶斯文本分类与马尔可夫网络图像去噪实现",
            font=("Arial", 12)
        )
        subtitle_label.pack()

        # 创建选项卡
        self.notebook = ttk.Notebook(self)
        self.notebook.pack(expand=True, fill=tk.BOTH, padx=10, pady=10)

        # 添加文本分类选项卡
        self.nb_tab = ttk.Frame(self.notebook)
        self.notebook.add(self.nb_tab, text="朴素贝叶斯文本分类")

        # 添加图像去噪选项卡
        self.mrf_tab = ttk.Frame(self.notebook)
        self.notebook.add(self.mrf_tab, text="马尔可夫网络图像去噪")

        # 添加关于选项卡
        self.about_tab = ttk.Frame(self.notebook)
        self.notebook.add(self.about_tab, text="关于")

        # 初始化各选项卡的内容
        self.setup_naive_bayes_tab()
        self.setup_mrf_tab()
        self.setup_about_tab()

        # 存储实验结果
        self.nb_results = None
        self.nb_visualizations = None
        self.mrf_results = None
        self.mrf_visualizations = None

        # 添加状态栏
        self.status_var = tk.StringVar()
        self.status_var.set("就绪")
        status_bar = ttk.Label(self, textvariable=self.status_var, relief=tk.SUNKEN, anchor=tk.W)
        status_bar.pack(side=tk.BOTTOM, fill=tk.X)

    def setup_naive_bayes_tab(self):
        """设置朴素贝叶斯选项卡的内容"""
        # 创建左右分栏
        left_frame = ttk.Frame(self.nb_tab)
        left_frame.pack(side=tk.LEFT, fill=tk.BOTH, expand=True, padx=5, pady=5)

        right_frame = ttk.Frame(self.nb_tab)
        right_frame.pack(side=tk.RIGHT, fill=tk.BOTH, expand=True, padx=5, pady=5)

        # 左侧:参数设置和运行按钮
        params_frame = ttk.LabelFrame(left_frame, text="参数设置")
        params_frame.pack(fill=tk.X, padx=5, pady=5)

        # 类别数量
        ttk.Label(params_frame, text="类别数量 (2-20):").grid(row=0, column=0, sticky=tk.W, padx=5, pady=5)
        self.categories_var = tk.IntVar(value=4)
        categories_spinbox = ttk.Spinbox(params_frame, from_=2, to=20, textvariable=self.categories_var, width=10)
        categories_spinbox.grid(row=0, column=1, sticky=tk.W, padx=5, pady=5)

        # 平滑参数alpha
        ttk.Label(params_frame, text="平滑参数 alpha:").grid(row=1, column=0, sticky=tk.W, padx=5, pady=5)
        self.alpha_var = tk.DoubleVar(value=1.0)
        alpha_spinbox = ttk.Spinbox(params_frame, from_=0.01, to=10.0, increment=0.1, textvariable=self.alpha_var,
                                    width=10)
        alpha_spinbox.grid(row=1, column=1, sticky=tk.W, padx=5, pady=5)

        # 启用集成模型
        self.ensemble_var = tk.BooleanVar(value=True)
        ensemble_check = ttk.Checkbutton(params_frame, text="启用集成模型", variable=self.ensemble_var)
        ensemble_check.grid(row=2, column=0, columnspan=2, sticky=tk.W, padx=5, pady=5)

        # 特征选项
        features_frame = ttk.LabelFrame(left_frame, text="特征设置")
        features_frame.pack(fill=tk.X, padx=5, pady=5)

        # N元语法范围
        ttk.Label(features_frame, text="N元语法范围:").grid(row=0, column=0, sticky=tk.W, padx=5, pady=5)
        self.ngram_var = tk.StringVar(value="1-2")
        ngram_combo = ttk.Combobox(features_frame, textvariable=self.ngram_var, values=["1-1", "1-2", "1-3"], width=8)
        ngram_combo.grid(row=0, column=1, sticky=tk.W, padx=5, pady=5)

        # 最大特征数
        ttk.Label(features_frame, text="最大特征数:").grid(row=1, column=0, sticky=tk.W, padx=5, pady=5)
        self.max_features_var = tk.IntVar(value=1000)
        max_features_spinbox = ttk.Spinbox(
            features_frame, from_=100, to=10000, increment=100, textvariable=self.max_features_var, width=10)
        max_features_spinbox.grid(row=1, column=1, sticky=tk.W, padx=5, pady=5)

        # 运行按钮
        run_button = ttk.Button(left_frame, text="运行文本分类", command=self.run_naive_bayes)
        run_button.pack(fill=tk.X, padx=5, pady=10)

        # 左侧:输出日志
        log_frame = ttk.LabelFrame(left_frame, text="运行日志")
        log_frame.pack(fill=tk.BOTH, expand=True, padx=5, pady=5)

        self.nb_log = scrolledtext.ScrolledText(log_frame, wrap=tk.WORD, height=10)
        self.nb_log.pack(fill=tk.BOTH, expand=True, padx=5, pady=5)
        self.nb_log.configure(state="disabled")

        # 右侧:结果展示区
        results_frame = ttk.LabelFrame(right_frame, text="实验结果")
        results_frame.pack(fill=tk.BOTH, expand=True, padx=5, pady=5)

        # 创建图表选择区域
        self.nb_view_buttons_frame = ttk.Frame(results_frame)
        self.nb_view_buttons_frame.pack(fill=tk.X, padx=5, pady=5)

        # 创建画布以显示图表
        self.nb_canvas_frame = ttk.Frame(results_frame)
        self.nb_canvas_frame.pack(fill=tk.BOTH, expand=True, padx=5, pady=5)

        # 默认显示提示信息
        ttk.Label(
            self.nb_canvas_frame,
            text="运行实验后此处将显示结果",
            font=("Arial", 14),
            anchor=tk.CENTER
        ).pack(expand=True)

    def setup_mrf_tab(self):
        """设置马尔可夫网络选项卡的内容"""
        # 创建左右分栏
        left_frame = ttk.Frame(self.mrf_tab)
        left_frame.pack(side=tk.LEFT, fill=tk.BOTH, expand=True, padx=5, pady=5)

        right_frame = ttk.Frame(self.mrf_tab)
        right_frame.pack(side=tk.RIGHT, fill=tk.BOTH, expand=True, padx=5, pady=5)

        # 左侧:参数设置和运行按钮
        params_frame = ttk.LabelFrame(left_frame, text="参数设置")
        params_frame.pack(fill=tk.X, padx=5, pady=5)

        # 图像类型选择
        ttk.Label(params_frame, text="图像模式:").grid(row=0, column=0, sticky=tk.W, padx=5, pady=5)
        self.pattern_var = tk.StringVar(value="text")
        pattern_combo = ttk.Combobox(
            params_frame, textvariable=self.pattern_var,
            values=["checkerboard", "horizontal_stripes", "vertical_stripes", "circle", "cross", "random", "text"],
            width=15
        )
        pattern_combo.grid(row=0, column=1, sticky=tk.W, padx=5, pady=5)

        # 图像尺寸
        ttk.Label(params_frame, text="图像尺寸:").grid(row=1, column=0, sticky=tk.W, padx=5, pady=5)
        self.size_var = tk.IntVar(value=100)
        size_spinbox = ttk.Spinbox(params_frame, from_=50, to=300, increment=10, textvariable=self.size_var, width=10)
        size_spinbox.grid(row=1, column=1, sticky=tk.W, padx=5, pady=5)

        # 噪声率
        ttk.Label(params_frame, text="噪声率:").grid(row=2, column=0, sticky=tk.W, padx=5, pady=5)
        self.noise_rate_var = tk.DoubleVar(value=0.1)
        noise_scale = ttk.Scale(params_frame, from_=0.01, to=0.5, orient=tk.HORIZONTAL,
                                variable=self.noise_rate_var, length=100)
        noise_scale.grid(row=2, column=1, sticky=tk.W, padx=5, pady=5)
        self.noise_label = ttk.Label(params_frame, text="0.10")
        self.noise_label.grid(row=2, column=2, sticky=tk.W)

        # 更新噪声率标签
        def update_noise_label(*args):
            self.noise_label.config(text=f"{self.noise_rate_var.get():.2f}")

        self.noise_rate_var.trace("w", update_noise_label)

        # 模型参数
        model_frame = ttk.LabelFrame(left_frame, text="模型参数")
        model_frame.pack(fill=tk.X, padx=5, pady=5)

        # 数据项权重eta
        ttk.Label(model_frame, text="数据项权重 eta:").grid(row=0, column=0, sticky=tk.W, padx=5, pady=5)
        self.eta_var = tk.DoubleVar(value=2.0)
        eta_spinbox = ttk.Spinbox(model_frame, from_=0.5, to=10.0, increment=0.5, textvariable=self.eta_var, width=10)
        eta_spinbox.grid(row=0, column=1, sticky=tk.W, padx=5, pady=5)

        # 光滑项权重beta
        ttk.Label(model_frame, text="光滑项权重 beta:").grid(row=1, column=0, sticky=tk.W, padx=5, pady=5)
        self.beta_var = tk.DoubleVar(value=1.0)
        beta_spinbox = ttk.Spinbox(model_frame, from_=0.1, to=5.0, increment=0.1, textvariable=self.beta_var, width=10)
        beta_spinbox.grid(row=1, column=1, sticky=tk.W, padx=5, pady=5)

        # 最大迭代次数
        ttk.Label(model_frame, text="最大迭代次数:").grid(row=2, column=0, sticky=tk.W, padx=5, pady=5)
        self.max_iters_var = tk.IntVar(value=10)
        iters_spinbox = ttk.Spinbox(model_frame, from_=1, to=50, increment=1, textvariable=self.max_iters_var, width=10)
        iters_spinbox.grid(row=2, column=1, sticky=tk.W, padx=5, pady=5)

        # 运行按钮
        run_button = ttk.Button(left_frame, text="运行图像去噪", command=self.run_mrf_denoising)
        run_button.pack(fill=tk.X, padx=5, pady=10)

        # 左侧:输出日志
        log_frame = ttk.LabelFrame(left_frame, text="运行日志")
        log_frame.pack(fill=tk.BOTH, expand=True, padx=5, pady=5)

        self.mrf_log = scrolledtext.ScrolledText(log_frame, wrap=tk.WORD, height=10)
        self.mrf_log.pack(fill=tk.BOTH, expand=True, padx=5, pady=5)
        self.mrf_log.configure(state="disabled")

        # 右侧:结果展示区
        self.mrf_results_frame = ttk.LabelFrame(right_frame, text="图像去噪结果")
        self.mrf_results_frame.pack(fill=tk.BOTH, expand=True, padx=5, pady=5)

        # 创建画布以显示图像和结果
        self.mrf_canvas_frame = ttk.Frame(self.mrf_results_frame)
        self.mrf_canvas_frame.pack(fill=tk.BOTH, expand=True, padx=5, pady=5)

        # 默认显示提示信息
        ttk.Label(
            self.mrf_canvas_frame,
            text="运行实验后此处将显示结果",
            font=("Arial", 14),
            anchor=tk.CENTER
        ).pack(expand=True)

        # 创建选项按钮,用于在不同结果之间切换
        self.mrf_view_buttons_frame = ttk.Frame(self.mrf_results_frame)
        self.mrf_view_buttons_frame.pack(fill=tk.X, padx=5, pady=5)

    def setup_about_tab(self):
        """设置关于选项卡的内容"""
        about_frame = ttk.Frame(self.about_tab)
        about_frame.pack(fill=tk.BOTH, expand=True, padx=20, pady=20)

        # 标题和版本信息
        ttk.Label(
            about_frame,
            text="Python机器学习:朴素贝叶斯文本分类与马尔可夫网络图像去噪实现",
            font=("Arial", 14, "bold")
        ).pack(pady=10)

        ttk.Label(
            about_frame,
            text="版本 1.0.0",
            font=("Arial", 12)
        ).pack(pady=5)

        # 分隔线
        ttk.Separator(about_frame, orient=tk.HORIZONTAL).pack(fill=tk.X, pady=15)

        # 项目说明
        ttk.Label(
            about_frame,
            text="项目说明",
            font=("Arial", 12, "bold")
        ).pack(anchor=tk.W, pady=5)

        description_text = """
本项目实现了两种经典的机器学习算法:

1. 朴素贝叶斯文本分类:
   - 提供自定义实现和sklearn实现的对比
   - 支持模型集成以提高性能
   - 包含性能可视化和参数优化功能

2. 马尔可夫随机场图像去噪:
   - 基于ICM算法的实现
   - 支持多种图像模式和参数设置
   - 提供实时去噪过程和结果可视化

本软件支持命令行和图形界面两种运行方式,使用matplotlib和tkinter进行可视化。
        """
        description = ttk.Label(about_frame, text=description_text, wraplength=600, justify=tk.LEFT)
        description.pack(fill=tk.X, pady=5)

        # 分隔线
        ttk.Separator(about_frame, orient=tk.HORIZONTAL).pack(fill=tk.X, pady=15)

        # 使用说明
        ttk.Label(
            about_frame,
            text="使用说明",
            font=("Arial", 12, "bold")
        ).pack(anchor=tk.W, pady=5)

        usage_text = """
1. 朴素贝叶斯文本分类:
   - 选择类别数量、平滑参数等
   - 点击"运行文本分类"按钮开始实验
   - 点击按钮查看不同的结果图表

2. 马尔可夫网络图像去噪:
   - 选择图像模式、噪声率和模型参数
   - 点击"运行图像去噪"按钮开始实验
   - 点击按钮查看不同阶段的图像和结果

3. 图表交互功能:
   - 点击按钮可以查看不同的结果图表
   - 点击"在新窗口中查看"按钮可以打开独立的图表窗口
   - 在独立窗口中可以使用matplotlib工具栏进行缩放、平移等操作
   - 可以将图表保存为PNG、JPG、SVG或PDF格式

命令行运行方式:
python ml_implementation.py --help
        """
        usage = ttk.Label(about_frame, text=usage_text, wraplength=600, justify=tk.LEFT)
        usage.pack(fill=tk.X, pady=5)

        # 分隔线
        ttk.Separator(about_frame, orient=tk.HORIZONTAL).pack(fill=tk.X, pady=15)

        # 版权信息
        ttk.Label(
            about_frame,
            text="版权所有 © 2023",
            font=("Arial", 10)
        ).pack(side=tk.BOTTOM, pady=10)

    def run_naive_bayes(self):
        """执行朴素贝叶斯文本分类实验"""
        # 清空日志和结果区域
        self.nb_log.configure(state="normal")
        self.nb_log.delete(1.0, tk.END)
        self.nb_log.configure(state="disabled")

        # 清空结果展示区域
        for widget in self.nb_canvas_frame.winfo_children():
            widget.destroy()

        # 清空按钮区域
        for widget in self.nb_view_buttons_frame.winfo_children():
            widget.destroy()

        # 创建进度指示
        progress_frame = ttk.Frame(self.nb_canvas_frame)
        progress_frame.pack(fill=tk.BOTH, expand=True)

        ttk.Label(progress_frame, text="正在运行朴素贝叶斯文本分类实验...", font=("Arial", 12)).pack(pady=20)

        progress_bar = ttk.Progressbar(progress_frame, mode="indeterminate")
        progress_bar.pack(fill=tk.X, padx=50, pady=10)
        progress_bar.start(10)

        # 获取参数
        try:
            categories = self.categories_var.get()
            alpha = self.alpha_var.get()
            ensemble = self.ensemble_var.get()

            # 设置全局参数
            args.categories = categories
            args.alpha = alpha
            args.ensemble = ensemble

            # 显示参数信息
            self.update_nb_log(f"开始运行朴素贝叶斯文本分类实验...\n")
            self.update_nb_log(f"参数设置:\n")
            self.update_nb_log(f"- 类别数量: {categories}\n")
            self.update_nb_log(f"- 平滑参数alpha: {alpha}\n")
            self.update_nb_log(f"- 启用集成模型: {'是' if ensemble else '否'}\n\n")

            # 更新状态栏
            self.status_var.set("正在运行朴素贝叶斯文本分类实验...")

            # 在单独的线程中运行实验以避免界面冻结
            thread = threading.Thread(target=self._run_nb_thread)
            thread.daemon = True
            thread.start()

        except Exception as e:
            self.update_nb_log(f"错误: {str(e)}\n")
            messagebox.showerror("错误", f"运行实验时出错: {str(e)}")

    def _run_nb_thread(self):
        """在单独线程中运行朴素贝叶斯实验"""
        try:
            # 运行实验
            self.nb_results, self.nb_visualizations = run_text_classification(self.update_nb_log)

            # 在主线程中显示结果
            self.after(100, self.show_nb_results)
        except Exception as e:
            # 在主线程中显示错误
            self.after(100, lambda: self.show_nb_error(str(e)))

    def show_nb_results(self):
        """显示朴素贝叶斯实验结果"""
        try:
            if not self.nb_results or not self.nb_visualizations:
                self.update_nb_log("实验未返回有效结果\n")
                return

            # 清空结果区域
            for widget in self.nb_canvas_frame.winfo_children():
                widget.destroy()

            # 清空按钮区域
            for widget in self.nb_view_buttons_frame.winfo_children():
                widget.destroy()

            # 创建预览区域
            preview_frame = ttk.Frame(self.nb_canvas_frame)
            preview_frame.pack(fill=tk.BOTH, expand=True)

            # 创建标题标签
            self.nb_preview_title = ttk.Label(
                preview_frame,
                text="点击下方按钮查看结果图表",
                font=("Arial", 12, "bold")
            )
            self.nb_preview_title.pack(pady=10)

            # 创建预览图像占位符
            self.nb_preview_canvas = ttk.Frame(preview_frame)
            self.nb_preview_canvas.pack(fill=tk.BOTH, expand=True)

            ttk.Label(self.nb_preview_canvas, text="选择要查看的图表类型").pack(expand=True)

            # 创建按钮区域
            button_frame = ttk.Frame(self.nb_view_buttons_frame)
            button_frame.pack(fill=tk.X)

            # 添加查看不同图表的按钮
            view_names = [name for name, _ in self.nb_visualizations]

            model_labels = {
                'confusion_matrix': '混淆矩阵',
                'accuracy_comparison': '准确率比较',
                'metrics_comparison': '性能指标比较',
                'class_distribution': '类别分布'
            }

            for i, (name, fig) in enumerate(self.nb_visualizations):
                btn_text = model_labels.get(name, name)
                btn = ttk.Button(
                    button_frame,
                    text=btn_text,
                    command=lambda n=name, f=fig: self.show_nb_preview(n, f)
                )
                btn.grid(row=0, column=i, padx=5, pady=5)

            # 添加查看指标的按钮
            ttk.Button(
                button_frame,
                text="查看性能指标",
                command=self.show_nb_metrics
            ).grid(row=0, column=len(view_names), padx=5, pady=5)

            # 显示第一个图表预览
            if view_names:
                self.show_nb_preview(view_names[0], self.nb_visualizations[0][1])

            # 更新状态栏
            self.status_var.set("朴素贝叶斯文本分类实验完成")

        except Exception as e:
            self.show_nb_error(str(e))

    def show_nb_preview(self, name, fig):
        """在预览区域显示图表"""
        # 更新标题
        model_labels = {
            'confusion_matrix': '混淆矩阵',
            'accuracy_comparison': '准确率比较',
            'metrics_comparison': '性能指标比较',
            'class_distribution': '类别分布'
        }
        self.nb_preview_title.config(text=model_labels.get(name, name))

        # 清空预览区域
        for widget in self.nb_preview_canvas.winfo_children():
            widget.destroy()

        # 创建Matplotlib图表预览
        canvas_frame = ttk.Frame(self.nb_preview_canvas)
        canvas_frame.pack(fill=tk.BOTH, expand=True)

        canvas = FigureCanvasTkAgg(fig, canvas_frame)
        canvas.draw()
        canvas.get_tk_widget().pack(fill=tk.BOTH, expand=True)

        # 添加"在新窗口中查看"按钮
        btn_frame = ttk.Frame(self.nb_preview_canvas)
        btn_frame.pack(fill=tk.X, pady=5)

        ttk.Button(
            btn_frame,
            text="在新窗口中查看",
            command=lambda: self.open_figure_window(fig, model_labels.get(name, name))
        ).pack(side=tk.RIGHT, padx=10)

    def show_nb_metrics(self):
        """显示朴素贝叶斯性能指标"""
        # 更新标题
        self.nb_preview_title.config(text="模型性能指标")

        # 清空预览区域
        for widget in self.nb_preview_canvas.winfo_children():
            widget.destroy()

        # 创建性能指标显示区域
        metrics_frame = ttk.Frame(self.nb_preview_canvas)
        metrics_frame.pack(fill=tk.BOTH, expand=True, padx=10, pady=10)

        # 获取指标
        scores = self.nb_results.get('scores', {})
        if not scores:
            ttk.Label(metrics_frame, text="未找到性能指标数据").pack(pady=20)
            return

        # 创建性能表格
        columns = ("accuracy", "precision", "recall", "f1")
        tree = ttk.Treeview(metrics_frame, columns=columns, show="headings")
        tree.heading("accuracy", text="准确率")
        tree.heading("precision", text="精确率")
        tree.heading("recall", text="召回率")
        tree.heading("f1", text="F1分数")

        # 设置列宽
        for col in columns:
            tree.column(col, width=100, anchor="center")

        # 添加数据
        model_labels = {
            'custom_nb': '自定义朴素贝叶斯',
            'sklearn_nb': 'Sklearn朴素贝叶斯',
            'ensemble': '集成模型',
            'best_nb': '最优参数模型'
        }

        for model_name, score in scores.items():
            display_name = model_labels.get(model_name, model_name)
            tree.insert("", tk.END, text=display_name, values=(
                f"{score.get('accuracy', 0):.4f}",
                f"{score.get('precision', 0):.4f}",
                f"{score.get('recall', 0):.4f}",
                f"{score.get('f1', 0):.4f}"
            ))

        # 添加滚动条
        scrollbar = ttk.Scrollbar(metrics_frame, orient="vertical", command=tree.yview)
        tree.configure(yscrollcommand=scrollbar.set)
        scrollbar.pack(side="right", fill="y")
        tree.pack(expand=True, fill="both")

        # 添加参数信息
        params_frame = ttk.LabelFrame(self.nb_preview_canvas, text="实验参数")
        params_frame.pack(fill=tk.X, padx=10, pady=10)

        ttk.Label(params_frame, text=f"类别数量: {args.categories}").grid(row=0, column=0, sticky="w", padx=10, pady=5)
        ttk.Label(params_frame, text=f"平滑参数alpha: {args.alpha}").grid(row=0, column=1, sticky="w", padx=10, pady=5)
        ttk.Label(params_frame, text=f"启用集成模型: {'是' if args.ensemble else '否'}").grid(row=1, column=0,
                                                                                              sticky="w", padx=10,
                                                                                              pady=5)

        # 如果有最佳参数信息,显示它
        if 'best_nb' in scores and 'best_alpha' in scores['best_nb']:
            ttk.Label(params_frame, text=f"最佳alpha参数: {scores['best_nb']['best_alpha']}").grid(row=1, column=1,
                                                                                                   sticky="w", padx=10,
                                                                                                   pady=5)

    def open_figure_window(self, fig, title="图形窗口"):
        """在新窗口中打开图表"""
        window = FigureWindow(self, fig, title)
        window.focus_set()

    def show_nb_error(self, error_msg):
        """显示朴素贝叶斯实验错误"""
        self.update_nb_log(f"错误: {error_msg}\n")
        messagebox.showerror("错误", f"运行实验时出错: {error_msg}")
        self.status_var.set("实验失败")

        # 清空结果区域
        for widget in self.nb_canvas_frame.winfo_children():
            widget.destroy()

        # 显示错误信息
        error_frame = ttk.Frame(self.nb_canvas_frame)
        error_frame.pack(fill=tk.BOTH, expand=True)

        ttk.Label(
            error_frame,
            text="实验运行失败",
            font=("Arial", 14, "bold"),
            foreground="red"
        ).pack(pady=10)

        ttk.Label(
            error_frame,
            text=error_msg,
            wraplength=400
        ).pack(pady=10)

    def run_mrf_denoising(self):
        """执行马尔可夫网络图像去噪实验"""
        # 清空日志和结果区域
        self.mrf_log.configure(state="normal")
        self.mrf_log.delete(1.0, tk.END)
        self.mrf_log.configure(state="disabled")

        # 清空结果展示区域
        for widget in self.mrf_canvas_frame.winfo_children():
            widget.destroy()

        # 清空按钮区域
        for widget in self.mrf_view_buttons_frame.winfo_children():
            widget.destroy()

        # 创建进度指示
        progress_frame = ttk.Frame(self.mrf_canvas_frame)
        progress_frame.pack(fill=tk.BOTH, expand=True)

        ttk.Label(progress_frame, text="正在运行马尔可夫网络图像去噪实验...", font=("Arial", 12)).pack(pady=20)

        progress_bar = ttk.Progressbar(progress_frame, mode="indeterminate")
        progress_bar.pack(fill=tk.X, padx=50, pady=10)
        progress_bar.start(10)

        # 获取参数
        try:
            pattern_type = self.pattern_var.get()
            size = self.size_var.get()
            noise_rate = self.noise_rate_var.get()
            eta = self.eta_var.get()
            beta = self.beta_var.get()
            max_iters = self.max_iters_var.get()

            # 设置全局参数
            args.noise_rate = noise_rate
            args.eta = eta
            args.beta = beta
            args.max_iters = max_iters

            # 显示参数信息
            self.update_mrf_log(f"开始运行马尔可夫网络图像去噪实验...\n")
            self.update_mrf_log(f"参数设置:\n")
            self.update_mrf_log(f"- 图像模式: {pattern_type}\n")
            self.update_mrf_log(f"- 图像尺寸: {size}\n")
            self.update_mrf_log(f"- 噪声率: {noise_rate}\n")
            self.update_mrf_log(f"- 数据项权重eta: {eta}\n")
            self.update_mrf_log(f"- 光滑项权重beta: {beta}\n")
            self.update_mrf_log(f"- 最大迭代次数: {max_iters}\n\n")

            # 更新状态栏
            self.status_var.set("正在运行马尔可夫网络图像去噪实验...")

            # 在单独的线程中运行实验以避免界面冻结
            thread = threading.Thread(target=self._run_mrf_thread)
            thread.daemon = True
            thread.start()

        except Exception as e:
            self.update_mrf_log(f"错误: {str(e)}\n")
            messagebox.showerror("错误", f"运行实验时出错: {str(e)}")

    def _run_mrf_thread(self):
        """在单独线程中运行马尔可夫网络实验"""
        try:
            # 运行实验
            self.mrf_results, self.mrf_visualizations = run_image_denoising(self.update_mrf_log)

            # 在主线程中显示结果
            self.after(100, self.show_mrf_results)
        except Exception as e:
            # 在主线程中显示错误
            self.after(100, lambda: self.show_mrf_error(str(e)))

    def update_mrf_log_and_preview(self, message, img=None, img_type=None):
        """更新马尔可夫网络实验日志和预览图像"""
        self.update_mrf_log(message + "\n")
        if img is not None and isinstance(img, np.ndarray):
            self.after(10, lambda: self._update_preview_image_main_thread(img, message, img_type))

    def _update_preview_image_main_thread(self, img, title="", img_type=None):
        """在主线程中更新预览图像"""
        try:
            # 将numpy数组转换为PIL图像,然后转换为PhotoImage
            plt.figure(figsize=(5, 5))
            plt.imshow(img, cmap='binary_r')
            if title:
                plt.title(title)
            plt.axis('off')

            # 保存为BytesIO对象
            buf = io.BytesIO()
            plt.savefig(buf, format='png')
            buf.seek(0)

            # 转换为PIL图像
            pil_img = Image.open(buf)

            # 调整大小以适应界面
            width, height = 300, 300
            pil_img = pil_img.resize((width, height), Image.LANCZOS)

            # 转换为PhotoImage
            tk_img = ImageTk.PhotoImage(pil_img)

            # 更新标签
            if hasattr(self, "mrf_image_label"):
                self.mrf_image_label.configure(image=tk_img)
                self.mrf_image_label.image = tk_img  # 保持引用以防止垃圾回收

            # 关闭图形,释放内存
            plt.close()

        except Exception as e:
            self.update_mrf_log(f"更新预览图像失败: {str(e)}\n")

    def show_mrf_results(self):
        """显示马尔可夫网络实验结果"""
        try:
            if not self.mrf_results or not self.mrf_visualizations:
                self.update_mrf_log("实验未返回有效结果\n")
                return

            # 清空结果区域
            for widget in self.mrf_canvas_frame.winfo_children():
                widget.destroy()

            # 清空按钮区域
            for widget in self.mrf_view_buttons_frame.winfo_children():
                widget.destroy()

            # 创建预览区域
            preview_frame = ttk.Frame(self.mrf_canvas_frame)
            preview_frame.pack(fill=tk.BOTH, expand=True)

            # 创建标题标签
            self.mrf_preview_title = ttk.Label(
                preview_frame,
                text="点击下方按钮查看结果图表",
                font=("Arial", 12, "bold")
            )
            self.mrf_preview_title.pack(pady=10)

            # 创建预览图像占位符
            self.mrf_preview_canvas = ttk.Frame(preview_frame)
            self.mrf_preview_canvas.pack(fill=tk.BOTH, expand=True)

            ttk.Label(self.mrf_preview_canvas, text="选择要查看的图表类型").pack(expand=True)

            # 创建按钮区域
            button_frame = ttk.Frame(self.mrf_view_buttons_frame)
            button_frame.pack(fill=tk.X)

            # 分类按钮,例如基本结果、迭代过程等
            button_categories = {
                "基本结果": ["original_image", "noisy_image", "denoised_image", "error_comparison"],
                "能量变化": ["energy_curve"],
                "迭代过程": [name for name, _ in self.mrf_visualizations if name.startswith("iteration_")],
                "参数比较": [name for name, _ in self.mrf_visualizations if
                             name.startswith("param_") or name == "parameter_comparison"]
            }

            # 翻译表
            translations = {
                "original_image": "原始图像",
                "noisy_image": "噪声图像",
                "denoised_image": "去噪结果",
                "error_comparison": "错误对比",
                "energy_curve": "能量曲线",
                "parameter_comparison": "参数比较"
            }

            # 创建类别标签框架
            categories_frame = ttk.LabelFrame(button_frame, text="结果分类")
            categories_frame.pack(fill=tk.X, padx=5, pady=5)

            # 为每个类别创建按钮行
            row = 0
            for category, names in button_categories.items():
                if names:
                    ttk.Label(categories_frame, text=f"{category}:").grid(row=row, column=0, sticky=tk.W, padx=5,
                                                                          pady=5)
                    col = 1
                    for name in names:
                        # 查找对应的图表
                        fig = None
                        for n, f in self.mrf_visualizations:
                            if n == name:
                                fig = f
                                break

                        if fig:
                            # 创建按钮
                            if name.startswith("iteration_"):
                                # 对于迭代图像,显示迭代号
                                iter_num = name.split("_")[1]
                                btn_text = f"迭代 {iter_num}"
                            elif name.startswith("param_"):
                                # 对于参数组合,提取eta和beta值
                                parts = name.split("_")
                                if len(parts) >= 3:
                                    eta = parts[1].replace("eta", "η=")
                                    beta = parts[2].replace("beta", "β=")
                                    btn_text = f"{eta},{beta}"
                                else:
                                    btn_text = name
                            else:
                                # 使用翻译
                                btn_text = translations.get(name, name)

                            btn = ttk.Button(
                                categories_frame,
                                text=btn_text,
                                width=12,
                                command=lambda n=name, f=fig: self.show_mrf_preview(n, f)
                            )
                            btn.grid(row=row, column=col, padx=3, pady=3)
                            col += 1

                            # 每行最多放置6个按钮
                            if col > 6:
                                row += 1
                                col = 1
                    row += 1

            # 添加指标按钮
            ttk.Button(
                categories_frame,
                text="性能指标",
                width=12,
                command=self.show_mrf_metrics
            ).grid(row=row, column=0, padx=3, pady=10)

            # 添加动画按钮
            ttk.Button(
                categories_frame,
                text="查看迭代动画",
                width=12,
                command=self.show_mrf_animation
            ).grid(row=row, column=1, padx=3, pady=10)

            # 显示原始图像预览
            self.show_mrf_preview("original_image", self.mrf_visualizations[0][1])

            # 更新状态栏
            self.status_var.set("马尔可夫网络图像去噪实验完成")

        except Exception as e:
            self.show_mrf_error(str(e))

    def show_mrf_preview(self, name, fig):
        """在预览区域显示图表"""
        # 更新标题
        translations = {
            "original_image": "原始图像",
            "noisy_image": "噪声图像",
            "denoised_image": "去噪结果",
            "error_comparison": "错误对比",
            "energy_curve": "能量曲线",
            "parameter_comparison": "参数组合比较"
        }

        if name.startswith("iteration_"):
            iter_num = name.split("_")[1]
            title = f"迭代 {iter_num} 结果"
        elif name.startswith("param_"):
            parts = name.split("_")
            if len(parts) >= 3:
                eta = parts[1].replace("eta", "η=")
                beta = parts[2].replace("beta", "β=")
                title = f"参数组合 {eta},{beta} 结果"
            else:
                title = name
        else:
            title = translations.get(name, name)

        self.mrf_preview_title.config(text=title)

        # 清空预览区域
        for widget in self.mrf_preview_canvas.winfo_children():
            widget.destroy()

        # 创建Matplotlib图表预览
        canvas_frame = ttk.Frame(self.mrf_preview_canvas)
        canvas_frame.pack(fill=tk.BOTH, expand=True)

        canvas = FigureCanvasTkAgg(fig, canvas_frame)
        canvas.draw()
        canvas.get_tk_widget().pack(fill=tk.BOTH, expand=True)

        # 添加"在新窗口中查看"按钮
        btn_frame = ttk.Frame(self.mrf_preview_canvas)
        btn_frame.pack(fill=tk.X, pady=5)

        ttk.Button(
            btn_frame,
            text="在新窗口中查看",
            command=lambda: self.open_figure_window(fig, title)
        ).pack(side=tk.RIGHT, padx=10)

    def show_mrf_metrics(self):
        """显示马尔可夫网络实验的性能指标"""
        # 更新标题
        self.mrf_preview_title.config(text="去噪性能指标")

        # 清空预览区域
        for widget in self.mrf_preview_canvas.winfo_children():
            widget.destroy()

        # 创建指标显示区域
        metrics_frame = ttk.Frame(self.mrf_preview_canvas)
        metrics_frame.pack(fill=tk.BOTH, expand=True, padx=10, pady=10)

        # 获取指标
        metrics = self.mrf_results.get('metrics', {})
        if not metrics:
            ttk.Label(metrics_frame, text="未找到性能指标数据").pack(pady=20)
            return

        # 创建指标表格
        tree = ttk.Treeview(metrics_frame, columns=("metric", "value"), show="headings")
        tree.heading("metric", text="指标名称")
        tree.heading("value", text="值")

        # 设置列宽
        tree.column("metric", width=200)
        tree.column("value", width=150)

        # 添加数据
        for name, value in metrics.items():
            if 'PSNR' in name:
                tree.insert("", tk.END, values=(name, f"{value:.2f} dB"))
            elif '百分比' in name:
                tree.insert("", tk.END, values=(name, f"{value:.2f}%"))
            else:
                tree.insert("", tk.END, values=(name, f"{value:.4f}"))

        # 添加滚动条
        scrollbar = ttk.Scrollbar(metrics_frame, orient="vertical", command=tree.yview)
        tree.configure(yscrollcommand=scrollbar.set)
        scrollbar.pack(side="right", fill="y")
        tree.pack(expand=True, fill="both")

        # 添加参数信息
        params_frame = ttk.LabelFrame(self.mrf_preview_canvas, text="实验参数")
        params_frame.pack(fill=tk.X, padx=10, pady=10)

        params = self.mrf_results.get('params', {})
        if params:
            row = 0
            col = 0
            for key, value in params.items():
                ttk.Label(params_frame, text=f"{key}: {value}").grid(row=row, column=col, sticky=tk.W, padx=10, pady=5)
                col += 1
                if col > 1:
                    col = 0
                    row += 1

    def show_mrf_animation(self):
        """显示马尔可夫网络迭代过程动画"""
        # 更新标题
        self.mrf_preview_title.config(text="迭代过程动画")

        # 清空预览区域
        for widget in self.mrf_preview_canvas.winfo_children():
            widget.destroy()

        # 创建动画窗口
        animation_window = tk.Toplevel(self)
        animation_window.title("马尔可夫网络图像去噪 - 迭代过程")
        animation_window.geometry("600x500")

        # 获取中间结果
        intermediate_results = self.mrf_results.get('intermediate_results', [])

        if not intermediate_results:
            ttk.Label(self.mrf_preview_canvas, text="未找到迭代过程数据").pack(pady=20)
            return

        # 在主窗口显示提示
        ttk.Label(
            self.mrf_preview_canvas,
            text="迭代过程动画已在新窗口打开",
            font=("Arial", 12)
        ).pack(pady=20)

        # 创建动画控制区域
        control_frame = ttk.Frame(animation_window)
        control_frame.pack(side=tk.BOTTOM, fill=tk.X, padx=10, pady=10)

        # 创建画布
        fig, ax = plt.subplots(figsize=(6, 6))
        canvas = FigureCanvasTkAgg(fig, master=animation_window)
        canvas.get_tk_widget().pack(fill=tk.BOTH, expand=True)

        # 显示初始图像
        ax.imshow(intermediate_results[0], cmap='binary_r')
        ax.set_title(f"初始噪声图像")
        ax.axis('off')
        canvas.draw()

        # 当前帧索引
        current_frame = tk.IntVar(value=0)

        # 更新帧的函数
        def update_frame(index):
            ax.clear()
            frame_idx = int(index)
            img = intermediate_results[frame_idx]
            ax.imshow(img, cmap='binary_r')
            if frame_idx == 0:
                ax.set_title(f"初始噪声图像")
            else:
                ax.set_title(f"迭代 {frame_idx}/{len(intermediate_results) - 1} 后的结果")
            ax.axis('off')
            canvas.draw()
            frame_label.config(text=f"帧 {frame_idx}/{len(intermediate_results) - 1}")

        # 播放/暂停控制
        playing = tk.BooleanVar(value=False)
        delay = 500  # 毫秒

        def toggle_play():
            playing.set(not playing.get())
            if playing.get():
                play_button.config(text="暂停")
                play_next_frame()
            else:
                play_button.config(text="播放")

        def play_next_frame():
            if playing.get():
                idx = current_frame.get() + 1
                if idx >= len(intermediate_results):
                    idx = 0
                current_frame.set(idx)
                slider.set(idx)
                update_frame(idx)
                animation_window.after(delay, play_next_frame)

        # 添加控件
        frame_label = ttk.Label(control_frame, text=f"帧 0/{len(intermediate_results) - 1}")
        frame_label.pack(side=tk.TOP)

        # 帧滑块
        slider = ttk.Scale(
            control_frame,
            from_=0,
            to=len(intermediate_results) - 1,
            orient=tk.HORIZONTAL,
            variable=current_frame,
            command=update_frame
        )
        slider.pack(side=tk.TOP, fill=tk.X, padx=10, pady=5)

        # 播放控制按钮
        button_frame = ttk.Frame(control_frame)
        button_frame.pack(side=tk.TOP, fill=tk.X)

        prev_button = ttk.Button(
            button_frame,
            text="上一帧",
            command=lambda: (current_frame.set(max(0, current_frame.get() - 1)),
                             slider.set(current_frame.get()),
                             update_frame(current_frame.get()))
        )
        prev_button.pack(side=tk.LEFT, padx=5)

        play_button = ttk.Button(
            button_frame,
            text="播放",
            command=toggle_play
        )
        play_button.pack(side=tk.LEFT, padx=5)

        next_button = ttk.Button(
            button_frame,
            text="下一帧",
            command=lambda: (current_frame.set(min(len(intermediate_results) - 1, current_frame.get() + 1)),
                             slider.set(current_frame.get()),
                             update_frame(current_frame.get()))
        )
        next_button.pack(side=tk.LEFT, padx=5)

        # 调整播放速度
        speed_frame = ttk.Frame(control_frame)
        speed_frame.pack(side=tk.TOP, fill=tk.X, pady=5)

        ttk.Label(speed_frame, text="播放速度:").pack(side=tk.LEFT)

        def update_speed(value):
            nonlocal delay
            speed_value = int(value)
            # 将1-10的速度值转换为毫秒延迟(1000ms到100ms)
            delay = int(1100 - speed_value * 100)
            speed_label.config(text=f"{speed_value}")

        speed_slider = ttk.Scale(
            speed_frame,
            from_=1,
            to=10,
            orient=tk.HORIZONTAL,
            command=update_speed
        )
        speed_slider.set(6)  # 默认速度
        speed_slider.pack(side=tk.LEFT, fill=tk.X, expand=True, padx=5)

        speed_label = ttk.Label(speed_frame, text="6")
        speed_label.pack(side=tk.LEFT)

        # 保存动画按钮
        save_button = ttk.Button(
            control_frame,
            text="保存动画",
            command=lambda: self.save_animation(intermediate_results)
        )
        save_button.pack(side=tk.BOTTOM, pady=10)

    def save_animation(self, frames):
        """保存动画为GIF文件"""
        try:
            from matplotlib.animation import FuncAnimation

            file_path = filedialog.asksaveasfilename(
                defaultextension=".gif",
                filetypes=[("GIF动画", "*.gif")]
            )

            if not file_path:
                return

            # 创建图形
            fig, ax = plt.subplots(figsize=(6, 6))

            # 初始化函数
            def init():
                ax.imshow(frames[0], cmap='binary_r')
                ax.set_title("初始噪声图像")
                ax.axis('off')
                return [ax]

            # 动画更新函数
            def update(frame_idx):
                ax.clear()
                img = frames[frame_idx]
                ax.imshow(img, cmap='binary_r')
                if frame_idx == 0:
                    ax.set_title(f"初始噪声图像")
                else:
                    ax.set_title(f"迭代 {frame_idx}/{len(frames) - 1}")
                ax.axis('off')
                return [ax]

            # 创建动画
            anim = FuncAnimation(
                fig,
                update,
                frames=len(frames),
                init_func=init,
                interval=200,
                blit=True
            )

            # 保存动画
            anim.save(file_path, writer='pillow', fps=5, dpi=100)

            messagebox.showinfo("保存成功", f"动画已保存到:\n{file_path}")

        except Exception as e:
            messagebox.showerror("保存失败", f"保存动画时出错:\n{str(e)}")

    def show_mrf_error(self, error_msg):
        """显示马尔可夫网络实验错误"""
        self.update_mrf_log(f"错误: {error_msg}\n")
        messagebox.showerror("错误", f"运行实验时出错: {error_msg}")
        self.status_var.set("实验失败")

        # 清空结果区域
        for widget in self.mrf_canvas_frame.winfo_children():
            widget.destroy()

        # 显示错误信息
        error_frame = ttk.Frame(self.mrf_canvas_frame)
        error_frame.pack(fill=tk.BOTH, expand=True)

        ttk.Label(
            error_frame,
            text="实验运行失败",
            font=("Arial", 14, "bold"),
            foreground="red"
        ).pack(pady=10)

        ttk.Label(
            error_frame,
            text=error_msg,
            wraplength=400
        ).pack(pady=10)

    def update_nb_log(self, message):
        """更新朴素贝叶斯实验日志"""
        self.nb_log.configure(state="normal")
        self.nb_log.insert(tk.END, message)
        self.nb_log.see(tk.END)
        self.nb_log.configure(state="disabled")
        self.update()

    def update_mrf_log(self, message):
        """更新马尔可夫网络实验日志"""
        self.mrf_log.configure(state="normal")
        self.mrf_log.insert(tk.END, message)
        self.mrf_log.see(tk.END)
        self.mrf_log.configure(state="disabled")
        self.update()


def main():
    """主函数"""
    app = App()
    app.mainloop()

    print("\n" + "=" * 80)
    print("Python机器学习:朴素贝叶斯文本分类与马尔可夫网络图像去噪实现")
    print("=" * 80 + "\n")

    # 选择任务
    task = args.task
    if args.interactive and not args.gui:
        print("可用任务:")
        print("1. 朴素贝叶斯文本分类")
        print("2. 马尔可夫网络图像去噪")
        print("3. 两者都运行")

        while True:
            try:
                choice = int(input("\n请选择要运行的任务(1-3): "))
                if 1 <= choice <= 3:
                    task = {1: 'nb', 2: 'mrf', 3: 'both'}[choice]
                    break
                print("无效选择,请输入1-3之间的数字!")
            except ValueError:
                print("请输入有效的数字!")
                task = 'both'  # 默认运行两个任务
                break

    results = {}

    if task == 'nb' or task == 'both':
        try:
            nb_results, _ = run_text_classification()
            results['naive_bayes'] = nb_results
        except Exception as e:
            logger.error(f"文本分类任务失败: {str(e)}")
            print(f"文本分类任务出错: {str(e)}")

    if task == 'mrf' or task == 'both':
        try:
            mrf_results, _ = run_image_denoising()
            results['mrf_denoising'] = mrf_results
        except Exception as e:
            logger.error(f"图像去噪任务失败: {str(e)}")
            print(f"图像去噪任务出错: {str(e)}")

    print("\n" + "=" * 80)
    print("任务完成!")
    print("=" * 80 + "\n")

    if args.save_results:
        print(f"结果已保存到 {args.output_dir} 目录")


if __name__ == "__main__":
    try:
        main()
    except Exception as e:
        logger.error(f"程序执行失败: {str(e)}", exc_info=True)
        print(f"程序执行出错: {str(e)}")
        print(f"请查看日志文件 {log_filename} 获取详细错误信息")

程序运行结果如下:

六、总结

本文探讨了无监督学习中数据分布建模的关键问题,重点介绍了两种概率图模型:贝叶斯网络和马尔可夫网络。通过模拟实验展示了它们在文本分类和图像去噪中的应用。

  1. 核心内容:
  • 贝叶斯网络采用有向图建模变量间的依赖关系,基于条件概率分解联合分布
  • 马尔可夫网络使用无向图表示对称依赖,通过团上的势函数定义概率分布
  • 实现朴素贝叶斯文本分类(准确率0.85-0.92)和MRF图像去噪(PSNR提升8-12dB)

     2.创新点:

  • 提出混合特征提取方法(TF-IDF结合n-gram)
  • 改进ICM算法实现并行优化
  • 开发交互式GUI可视化分析工具

      3.应用价值:

  • 文本分类准确率较基线提升5-8%
  • 图像去噪错误率降低60-80%
  • 为复杂数据建模提供有效解决方案

该研究为理解数据分布和特征关系提供了实用框架,代码实现已开源,便于学术和工业应用。

Logo

脑启社区是一个专注类脑智能领域的开发者社区。欢迎加入社区,共建类脑智能生态。社区为开发者提供了丰富的开源类脑工具软件、类脑算法模型及数据集、类脑知识库、类脑技术培训课程以及类脑应用案例等资源。

更多推荐