# 配套文章：03-training-and-generation.md
# 本轮新增示例未执行；环境和运行方式见上级 GUIDE.md。

import torch
from torch import nn

# 用零初始化让第一次分布易于手算：三个候选的概率均为 1/3。
# 二维输入是教学特征，不是人为定义的真实语义坐标。
features = torch.tensor([[1.0, 0.0]])
layer = nn.Linear(2, 3, bias=False)
nn.init.zeros_(layer.weight)
optimizer = torch.optim.SGD(layer.parameters(), lr=0.3)
target = torch.tensor([0], dtype=torch.long)  # 第 0 类是本例正确目标“商家”。

logits = layer(features)
# CrossEntropyLoss 内部处理 log-softmax，输入应为原始 logits，不先做 softmax。
loss = nn.CrossEntropyLoss()(logits, target)
print("更新前概率：", logits.detach().softmax(dim=-1))
print("更新前损失：", loss.item())

# 清理旧梯度 → 反向传播 → 参数更新；本例只有一次更新，不宣称已经收敛。
optimizer.zero_grad()
loss.backward()
print("权重梯度：", layer.weight.grad)
optimizer.step()
with torch.no_grad():
    print("更新后概率：", layer(features).softmax(dim=-1))
