You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

基于鸢尾花(Iris)数据集的预测及基础分类代码技术咨询

完成鸢尾花数据集分类与预测的完整实现

Hey there! Let's wrap up your Iris classification code and get it making reliable predictions. I'll walk through the full workflow step by step, fixing your incomplete code and adding all necessary pieces for training and inference.

1. 补全数据集读取代码

First, let's finish loading the test dataset properly—you had a truncation in your original code:

import tensorflow as tf
import pandas as pd
COLUMN_NAMES = [ 'SepalLength', 'SepalWidth', 'PetalLength', 'PetalWidth', 'Species' ]

# 导入训练数据集
training_dataset = pd.read_csv('iris_training.csv', names=COLUMN_NAMES, header=0)
train_x = training_dataset.iloc[:, 0:4]  # 提取前4个特征列
train_y = training_dataset.iloc[:, 4]    # 提取标签列

# 导入测试数据集(补全代码)
test_dataset = pd.read_csv('iris_test.csv', names=COLUMN_NAMES, header=0)
test_x = test_dataset.iloc[:, 0:4]
test_y = test_dataset.iloc[:, 4]

2. 预处理标签数据

Since we're doing multi-class classification, we need to convert integer labels (0, 1, 2 for Iris setosa, versicolor, virginica) into one-hot encoded vectors. This helps the model learn more effectively:

# 转换为独热编码,适配多分类任务
train_y_onehot = tf.keras.utils.to_categorical(train_y, num_classes=3)
test_y_onehot = tf.keras.utils.to_categorical(test_y, num_classes=3)

3. 构建分类模型

The Iris dataset is small and straightforward, so a shallow neural network will work perfectly. Here's a simple but effective model:

model = tf.keras.Sequential([
    tf.keras.layers.Dense(16, activation='relu', input_shape=(4,)),  # 输入层+第一层隐藏层,接收4个特征
    tf.keras.layers.Dense(8, activation='relu'),                     # 第二层隐藏层
    tf.keras.layers.Dense(3, activation='softmax')                   # 输出层,3个类别用softmax输出概率
])

# 编译模型,配置优化器、损失函数和评估指标
model.compile(optimizer='adam',
              loss='categorical_crossentropy',
              metrics=['accuracy'])

4. 训练模型

Now let's train the model on our training data, using the test set to monitor performance during training:

history = model.fit(train_x, train_y_onehot,
                    epochs=50,        # 训练轮次
                    batch_size=8,     # 每批次样本数
                    validation_data=(test_x, test_y_onehot))  # 验证集

5. 评估模型性能

After training, let's check how well the model performs on unseen test data:

test_loss, test_acc = model.evaluate(test_x, test_y_onehot)
print(f'Test set accuracy: {test_acc:.2f}')

6. 进行预测

Once the model is trained, you can use it to predict on new or existing data. Here's how to make predictions for individual samples or the entire test set:

# 对整个测试集进行预测
predictions = model.predict(test_x)

# 查看第一个样本的预测结果(取概率最高的类别)
sample_index = 0
predicted_class = tf.argmax(predictions[sample_index]).numpy()
actual_class = test_y.iloc[sample_index]
print(f'Predicted class: {predicted_class}, Actual class: {actual_class}')

# 预测单个新样本(示例输入一组鸢尾花特征)
new_sample = [[5.1, 3.5, 1.4, 0.2]]  # 这是Iris setosa的典型特征值
new_prediction = model.predict(new_sample)
predicted_species_idx = tf.argmax(new_prediction).numpy()
species_names = ['Setosa', 'Versicolor', 'Virginica']
print(f'Predicted species: {species_names[predicted_species_idx]}')

额外小提示

  • If your accuracy isn't as high as expected, try adjusting the number of epochs, batch size, or adding one more hidden layer.
  • Ensure your iris_training.csv and iris_test.csv files are in the same directory as your script, or provide the full file path when loading.
  • You can plot the training history (loss and accuracy over epochs) using Matplotlib to visualize how the model improves over time.

内容的提问来源于stack exchange,提问作者Kyrylo Kalashnikov

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.25 04:23:26