用mindspore 2.0.0版本实现“朴素贝叶斯垃圾邮件分类”,在运行过程中出现报错。
收藏回复举报
用mindspore 2.0.0版本实现“朴素贝叶斯垃圾邮件分类”,在运行过程中出现报错。
t('forum.solved') 已解决
发表于2023-06-02 09:46:23
0 查看
数据在rar压缩包中,数据和运行的代码在同一级目录下

原本的数据类型是data类型的数据,但是data类型的数据,没有找到加载到程序里面的方法,所以将数据先转成为CSV类型,再将CSV类型的数据加载到程序中。

然后代码定义了一个朴素贝叶斯分类器类。该类有两个参数:垃圾邮件的先验概率和非垃圾邮件的先验概率。该课程还有两本词典:一本用于垃圾词,另一本用于非垃圾词。字典分别存储每个单词在垃圾邮件和非垃圾邮件中出现的次数。

该类还有一个 forward() 方法。 forward() 方法计算数据在每个假设下的对数似然。垃圾邮件的对数似然是通过将垃圾邮件的先验概率乘以垃圾邮件字数的 log-softmax 来计算的。 ham 的对数似然是以类似的方式计算的。

然后 forward() 方法返回具有最高对数似然的类。

然后代码使用 Adam 优化器训练模型。 Adam 优化器是一种使用的随机梯度下降优化器。该代码训练模型 10 个时期。

训练模型后,代码会在 spambase 数据集上评估模型。代码通过计算模型正确分类电子邮件的次数来计算模型的准确性。然后代码打印模型的准确性。

代码如下:

import mindspore
import mindspore.nn as nn
import mindspore.dataset as ds
import pandas as pd
#将data类型的数据集转换为csv类型
data = pd.read_table("./spambase.data")
print(data)  # 打印数据

data.to_csv('./spambase.csv',sep='|',index=False)  # data转成csv

# 读取生成的machine.csv文件进行验证
csv = pd.read_csv('./spambase.csv')
print(csv.head(15)) # 打印前15行
# Load the data set
data_set = ds.CSVDataset("./spambase.csv", num_parallel_workers=1)
data_set = data_set.shuffle(buffer_size=1024)
data_set = data_set.batch(batch_size=128)

# Define the model
class NaiveBayes(nn.Cell):
    def __init__(self):
        super().__init__()

        # Define the prior probabilities
        self.spam_prior = 0.5
        self.ham_prior = 0.5

        # Define the vocabulary
        self.vocabulary = set()
        for data in data_set:
            print(data)
            for word in data:
                print(word)
                self.vocabulary.add(word)

        # Define the word counts
        self.spam_word_counts = {}
        self.ham_word_counts = {}
        for data in data_set:
            for word in data:
                if word not in self.spam_word_counts:
                    self.spam_word_counts[word] = 0
                if word not in self.ham_word_counts:
                    self.ham_word_counts[word] = 0
                self.spam_word_counts[word] += data[-1]
                self.ham_word_counts[word] += data[-1]

        # Define the model parameters
        self.spam_theta = nn.Parameter(mindspore.Tensor(self.spam_word_counts, dtype=mindspore.float32))
        self.ham_theta = nn.Parameter(mindspore.Tensor(self.ham_word_counts, dtype=mindspore.float32))

    def forward(self, x):
        # Calculate the log-likelihood of the data under each hypothesis
        spam_log_likelihood = self.spam_prior * nn.logsoftmax(self.spam_theta, axis=1)
        ham_log_likelihood = self.ham_prior * nn.logsoftmax(self.ham_theta, axis=1)

        # Return the class with the highest log-likelihood
        return mindspore.argmax(spam_log_likelihood + ham_log_likelihood, axis=1)

# Train the model
model = NaiveBayes()
optimizer = mindspore.optimizer.Adam(learning_rate=0.001)
for epoch in range(10):
    for data in data_set:
        # Convert the strings to numbers
        data = [str2float(word) for word in data]
        loss = model(data[:-1])
        optimizer.zero_grad()
        loss.backward()
        optimizer.step()

# Evaluate the model
accuracy = 0
for data in data_set:
    # Convert the strings to numbers
    data = [str2float(word) for word in data]
    prediction = model(data[:-1])
    if prediction == data[-1]:
        accuracy += 1
accuracy /= len(data_set)
print("Accuracy:", accuracy)

代码出现的报错如下:

[ERROR] MD(,516c,?):2023-6-2 9:21:22 [mindspore\ccsrc\minddata\dataset\engine\datasetops\source\csv_op.cc:744] mindspore::dataset::CsvOp::ColMapAnalyse] Invalid parameter, duplicate column 0.64 for csv file: ./spambase.csv
[ERROR] MD(,516c,?):2023-6-2 9:21:22 [mindspore\ccsrc\minddata\dataset\engine\datasetops\source\csv_op.cc:695] mindspore::dataset::CsvOp::ComputeColMap] Invalid file, failed to get column name list from csv file: ./spambase.csv
Traceback (most recent call last):
  line 57, in <module>
    model = NaiveBayes()
  line 24, in __init__
    for data in data_set:
RuntimeError: Exception thrown from dataset pipeline. Refer to 'Dataset Pipeline Error Message'. 

------------------------------------------------------------------
- Dataset Pipeline Error Message: 
------------------------------------------------------------------
[ERROR] Invalid file, failed to get column name list from csv file: ./spambase.csv.

------------------------------------------------------------------
- C++ Call Stack: (For framework developers) 
------------------------------------------------------------------
mindspore\ccsrc\minddata\dataset\engine\datasetops\source\csv_op.cc(696).

此外

str2float
这个函数飘红,但是代码没有运行到此处,还没有在这里报错

我要发帖子