week2

目标:

尝试完成一个多分类任务的训练:一个随机向量,哪一维数字最大就属于第几类。

内容:

python 复制代码
import torch
import torch.nn as nn
import torch.optim as optim

# 1. 固定随机种子
torch.manual_seed(42)

# 2. 定义参数
input_dim = 5
num_classes = 5
num_samples = 10000
epochs = 100
batch_size = 128

# 3. 生成训练数据
X_train = torch.rand(num_samples, input_dim)

# 找出每个向量中最大值的位置
y_train = torch.argmax(X_train, dim=1)

# 生成独立测试数据
X_test = torch.rand(2000, input_dim)
y_test = torch.argmax(X_test, dim=1)

# 4. 定义神经网络
class Model(nn.Module):
    def __init__(self):
        super().__init__()

        self.fc = nn.Linear(5, 5)

    def forward(self, x):
        return self.fc(x)

model = Model()

# 5. 定义损失函数和优化器
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=0.01)

# 6. 训练模型
for epoch in range(epochs):
    model.train()

    indices = torch.randperm(num_samples)
    total_loss = 0
    total_correct = 0

    for start in range(0, num_samples, batch_size):
        idx = indices[start:start + batch_size]

        X_batch = X_train[idx]
        y_batch = y_train[idx]

        # 前向传播
        outputs = model(X_batch)

        # 计算损失
        loss = criterion(outputs, y_batch)

        # 梯度清零
        optimizer.zero_grad()

        # 反向传播
        loss.backward()

        # 更新参数
        optimizer.step()

        total_loss += loss.item() * len(idx)

        preds = torch.argmax(outputs, dim=1)
        total_correct += (preds == y_batch).sum().item()

    if (epoch + 1) % 10 == 0:
        print(
            f"Epoch {epoch + 1}, "
            f"Loss: {total_loss / num_samples:.4f}, "
            f"Accuracy: {total_correct / num_samples:.2%}"
        )

# 7. 测试模型
model.eval()

with torch.no_grad():
    outputs = model(X_test)
    predictions = torch.argmax(outputs, dim=1)

    accuracy = (predictions == y_test).float().mean()

    print(f"\n测试准确率: {accuracy.item():.2%}")

# 8. 预测新的随机向量
x = torch.tensor([
    [0.12, 0.35, 0.91, 0.43, 0.28]
])

with torch.no_grad():
    output = model(x)
    prediction = torch.argmax(output, dim=1)

print(f"预测类别:第 {prediction.item() + 1} 类")
print(f"真实类别:第 {torch.argmax(x).item() + 1} 类")
相关推荐
yangmu32031 小时前
Codex 提示词优化教程:用四要素把模糊需求变成可执行任务
人工智能·学习
海盗12341 小时前
AI 新闻日报 2026-10-01:OpenAI 把 ChatGPT 变成智能体平台,VS Code 上线多模型编排,国产算力补到内核层
人工智能·chatgpt·机器人·人工智能aigc
龙亘川2 小时前
数字化赋能基层协同治理:亘川智城一网统管平台落地实践思考
大数据·人工智能·智慧城市·开源软件·数据可视化
坏小虎2 小时前
Codex 中的 Current checkout 和 New worktree 怎么选?
人工智能
陈天伟教授2 小时前
DeepSeek Harness 生态里的学术写作插件
人工智能·自然语言处理
浪子明X2 小时前
注意力不是全连接层换名字:多头自注意力的张量实验
人工智能
打工仔折腾 AI2 小时前
从Attention到BERT:双向预训练语言模型到底解决了什么问题
人工智能·后端·python·深度学习·语言模型·bert
正经教主2 小时前
【FDE系列】阶段3:Day 56:评测体系入门 — 建立你的黄金评测集
人工智能·fde
鲲穹AI种草2 小时前
自媒体矩阵批量剪辑怎么选?鲲剪短视频批量处理工具横向评测
人工智能·音视频·媒体