DAY 35


一、模型结构可视化

理解模型的两大关键:损失函数如何定义任务 ,以及层设计如何决定参数总量

1. nn.Module 自带方法
  • 打印模型结构print(model) 直接输出各层名称与顺序。
  • 查看参数详情model.named_parameters() 可遍历所有可训练参数,获取名称、形状(param.shape)。
  • 权重分布分析 :提取 weight 参数并转为 numpy 数组,绘制直方图观察各层权重的均值、标准差、最值等统计量。这些信息可辅助调整学习率、初始化方法或正则化强度。
2. torchsummary 库的 summary
  • 安装:pip install torchsummary
  • 用法:summary(model, input_size=(4,))
  • 必须提供 input_size(不含批量维度),因为 PyTorch 是动态图,需通过虚拟输入前向传播一次来推断各层的输出形状和参数量。
  • input_size 格式取决于模型类型:MLP 用 (特征数,),CNN 用 (通道, 高, 宽),RNN 用 (序列长度, 特征数)
  • 报告中参数量的计算:如 Linear(4,10) 有 4×10 个权重 + 10 个偏置 = 50 个参数;Linear(10,3) 有 10×3+3=33,总参数量 83。
3. torchinfo 库的 summary
  • 提供比 torchsummary 更详细的摘要,包括每层输入/输出形状、参数量、计算量等,常与 TensorBoard 配合使用。
  • 用法相同:summary(model, input_size=(4,))

二、进度条功能(tqdm)

tqdm 用于在循环中可视化训练进度。

1. 手动更新
python 复制代码
from tqdm import tqdm
with tqdm(total=10, desc="训练进度", unit="epoch") as pbar:
    for i in range(10):
        # 训练代码...
        pbar.update(1)          # 手动前进1步
        pbar.set_postfix({'Loss': f'{loss:.4f}'})  # 显示实时指标
  • desc:左侧描述文字;unit:进度单位(如 epochbatch)。
  • set_postfix 可在右侧动态显示字典形式的附加信息。
2. 自动更新
python 复制代码
for i in tqdm(range(10), desc="处理任务", unit="个"):
    # 循环体,无需手动update

直接将可迭代对象传入 tqdm(),自动完成进度更新。

在训练代码中的应用:创建总步数为 num_epochs 的进度条,每完成一定轮次后调用 pbar.update(n) 并刷新 postfix 显示损失值,最终确保进度条到达 100%。


三、模型推理(评估)

训练完成后在测试集上评估模型性能。

关键步骤
  1. 评估模式model.eval() 固定 Dropout、BatchNorm 等训练特定层,保证输出稳定。

    • 注意:此模式不自动关闭梯度计算,仍需手动处理。
  2. 禁用梯度with torch.no_grad(): 避免构建计算图,节省显存并加速推理。

  3. 预测与准确率

    python 复制代码
    outputs = model(X_test)
    _, predicted = torch.max(outputs, 1)          # 取每行最大值的索引作为类别
    correct = (predicted == y_test).sum().item()  # 正确个数
    accuracy = correct / y_test.size(0)           # 准确率
    • torch.max(outputs, dim=1) 返回最大值和索引,_ 忽略最大值,仅获取预测类别。
    • 所有操作在 GPU 张量上完成,无需转换到 CPU 或使用 sklearn,以保持高效。
知识串联
  • 训练阶段:定义损失(CrossEntropyLoss)与优化器(SGD),在循环中进行前向、反向和参数更新。
  • 推理阶段:模型参数固定,仅需前向传播,通过 eval() + no_grad() 确保高效、稳定的输出。
相关推荐
YCOSA20256 小时前
系统工具使用检测.exe (仅供娱乐)
python·microsoft
2601_962295336 小时前
如何python实现网页的自动化
python·selenium·beautifulsoup·requests·网页自动化
ofoxcoding6 小时前
GPT Image 2.5 API 实战:Python 调用实现图片生成与编辑
人工智能·python·gpt·ai
张小姐的猫7 小时前
【AI大模型接入SDK】 —— Gemini接入封装
android·数据结构·数据库·c++·人工智能·python
Elastic 中国社区官方博客7 小时前
Elasticsearch Python DSL 客户端开发
大数据·数据库·python·elasticsearch·搜索引擎·全文检索
计算机编程-吉哥8 小时前
脑肿瘤MRI智能识别系统:基于深度学习的像素级脑肿瘤语义分割平台【计算机毕业设计选题推荐】
人工智能·python·深度学习·算法·毕业设计·课程设计·大数据毕业设计选题推荐
小木_.8 小时前
Python 离线识别滑块缺口距离,项目推荐
开发语言·python·滑块识别·人机验证·滑块缺口·缺口识别
Cenxi8 小时前
Python字符串方法练习手册
人工智能·python
liliangcsdn8 小时前
因子权重矩阵处理-因子权重收缩Shrinkage算法的探索
开发语言·python·机器学习
用户8356290780519 小时前
使用 Python 在 Excel 中添加和编辑形状
后端·python