新闻详情

新闻详情

首页 / 资讯中心 / 详情

一个模板搞定四类视觉模型:CNN与ViT的统一训练部署

发布时间:2026/9/26 11:30:47来源:尧图网络
一个模板搞定四类视觉模型:CNN与ViT的统一训练部署
这次我们来看一份能直接套用的深度学习训练模板。主题是 CNN但又不只是 CNN——用同一套 Python 代码把 AlexNet、VGG、ResNet、ViT 四类主流视觉模型的训练、验证、导出和部署全流程串起来。很多新手刚接触图像分类时问题往往不是“模型不会写”而是“训练脚本不知道怎么组织换一个模型就要重写一遍”。这份模板要解决的就是这么个问题一份代码通过改一个参数切换模型从数据读取一直跑到接口部署。最值得关注的几个点支持四类架构统一训练自动适配自定义分类数带完整的训练日志、模型保存、验证评估可导出 ONNX 做推理加速还能用 Flask 包一个轻量 API 出来。硬件门槛按常见深度学习环境来CPU 可以跑小数据集GPU 更适合完整训练显存占用取决于模型、分辨率和 batch size。下面会按“环境准备 → 代码结构 → 数据组织 → 训练验证 → 部署 API → 性能排错”的顺序把整个流程完整过一遍。如果你正准备从单模型脚本走向工程化训练这篇文章可以直接收藏。1. 核心能力速览能力项说明项目类型深度学习视觉分类训练与部署模板支持模型AlexNet、VGG16、ResNet50、ViT-B/16框架依赖PyTorch、TorchVision、Pillow、tqdm、Flask输入数据ImageFolder 目录结构支持自定义类别训练能力带验证、早停、学习率调整、模型保存部署能力模型导出 ONNXFlask API 调用批量任务支持批量图片推理配合循环或队列实现硬件要求CPU 可跑GPU 加速推荐显存随模型结构变化启动方式Python 脚本训练flask run或直接执行app.py启动 API接口 API提供/predict接口返回类别与置信度适合场景课程实验、毕业设计、小规模工程落地前的原型验证表格里的“显存占用”“训练耗时”没有写死因为这些和数据集规模、图像分辨率、batch size、是否使用预训练权重都有关系。更稳妥的做法是用一份小的子集先跑通再逐步放大。2. 为什么需要一份统一训练模板AlexNet 是 CNN 的奠基之作结构简单适合理解卷积、池化、全连接的基本流程。VGG 用连续小卷积核替代大卷积核探索了深度与性能的关系但参数量大、显存开销不低。ResNet 引入残差连接解决深层网络退化问题是目前工程中使用最广泛的主干网络之一。ViT 属于 Transformer 架构在视觉上的代表它不依赖卷积而是把图像切成 patch 后做自注意力计算与 CNN 形成完全不同的技术路线。很多初学者会把这四个模型分开写四份脚本。这样做的问题是数据预处理逻辑重复、训练循环冗余、模型保存格式不统一、后期部署时每个模型都要单独适配。一旦模型切换还要重新调试。统一模板的价值在于把数据加载、训练循环、验证评估、模型导出这些“不会经常变的部分”固定下来把模型结构这种“会变的模块”用配置项隔离。你真正需要关心的只有数据和模型参数而不是反复复制粘贴训练代码。从学习角度看统一模板也有利于做横向对比。同样一份数据集分别跑 AlexNet、VGG16、ResNet50、ViT你观察到的训练曲线、显存占用、收敛速度差异就是这四类架构最直观的技术特征。这种对比比单独看某篇论文的插图更有体感。3. 适用场景与使用边界这份模板适合这些场景图像分类实验比如手写数字、猫狗分类、工业零件缺陷分类模型架构对比在自定义数据集上微调预训练模型把训练好的模型快速封装成 HTTP 接口供前端或移动端调用。它不适合的场景也很明确目标检测、实例分割、OCR、人脸识别等复杂任务需要专门的框架和模型头模板没有覆盖大规模分布式训练需要扩展 DDP 或 DeepSpeed生产级高并发服务需要更完善的鉴权、限流和监控Flask Demo 只是原型。版权和合规方面要特别注意。预训练权重多数来自 ImageNet 等数据集使用时要查看对应许可证。如果数据集包含人脸、车牌、医疗影像或他人作品必须获得授权并做脱敏处理避免侵犯隐私和版权。部署到公网时接口要加访问控制不能把内网服务直接裸奔暴露。4. 环境准备与前置条件推荐环境操作系统Windows 10/11、Ubuntu 20.04 或 macOSCPU 运行Python 3.8 或更高版本PyTorch 2.xCPU 版或 CUDA 版均可NVIDIA GPU 时确认驱动支持 CUDA 11.8/12.x磁盘空间代码和依赖约 10GBImageNet 级数据集另算安装核心依赖pip install torch torchvision pillow tqdm flask如果需要导出 ONNX再安装pip install onnx onnxruntime安装完成后用一段短代码确认环境import torch print(PyTorch:, torch.__version__) print(CUDA available:, torch.cuda.is_available()) if torch.cuda.is_available(): print(GPU:, torch.cuda.get_device_name(0))如果torch.cuda.is_available()返回False需要去 PyTorch 官网选择对应 CUDA 版本重新安装或检查驱动。5. 项目结构与配置定义建议按下面这种结构组织项目目录cnn_train_template/ ├── data/ │ ├── train/ │ │ ├── class0/ │ │ │ ├── img001.jpg │ │ │ └── img002.jpg │ │ └── class1/ │ │ └── ... │ └── val/ │ ├── class0/ │ └── class1/ ├── checkpoint/ ├── runs/ ├── train.py ├── config.py ├── model.py ├── data_utils.py ├── export_onnx.py └── app.py训练脚本统一从config.py读取参数。这样你不需要为了一个 batch size 去改业务代码。config.py示例import torch class Config: # 数据路径 train_dir data/train val_dir data/val # 模型参数 model_name resnet50 # alexnet / vgg16 / resnet50 / vit num_classes 10 pretrained True input_size 224 # 训练参数 epochs 30 batch_size 16 lr 1e-3 weight_decay 1e-4 num_workers 2 device cuda if torch.cuda.is_available() else cpu # 早停与保存 patience 5 save_path checkpoint/best_model.pth # 推理部署 onnx_path checkpoint/best_model.onnx api_host 127.0.0.1 api_port 5000这里input_size默认为 224是因为 ResNet、ViT 等主流模型预设输入都是 224x224AlexNet 原始输入也是 224x224VGG 同样是 224x224。如果你的数据分辨率不一致需要在预处理里统一缩放。6. 数据加载与预处理先用 TorchVision 的ImageFolder加载目录数据再对训练集做随机裁剪、随机水平翻转、归一化。验证集只做缩放和中心裁剪不做增强。data_utils.py代码import torch from torchvision import datasets, transforms def get_transforms(input_size224): train_transform transforms.Compose([ transforms.RandomResizedCrop(input_size), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) val_transform transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(input_size), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) return train_transform, val_transform def build_loaders(cfg): train_transform, val_transform get_transforms(cfg.input_size) train_dataset datasets.ImageFolder(cfg.train_dir, transformtrain_transform) val_dataset datasets.ImageFolder(cfg.val_dir, transformval_transform) train_loader torch.utils.data.DataLoader( train_dataset, batch_sizecfg.batch_size, shuffleTrue, num_workerscfg.num_workers, pin_memoryTrue ) val_loader torch.utils.data.DataLoader( val_dataset, batch_sizecfg.batch_size, shuffleFalse, num_workerscfg.num_workers, pin_memoryTrue ) return train_loader, val_loader, val_dataset.classes注意ImageFolder要求目录下每个子文件夹代表一个类别文件夹名就是类别名。如果你的数据是 CSV 标注格式需要先整理成目录结构或者自己写一个Dataset子类。对入门阶段来说ImageFolder是最省事的。7. 模型定义与选择在model.py里写一个统一的模型工厂函数。根据cfg.model_name返回对应模型并自动替换最后一层分类头。import torch.nn as nn from torchvision import models def build_model(name, num_classes, pretrainedTrue): name name.lower() weights None if name alexnet: if pretrained: weights models.AlexNet_Weights.IMAGENET1K_V1 model models.alexnet(weightsweights) in_feat model.classifier[-1].in_features model.classifier[-1] nn.Linear(in_feat, num_classes) elif name vgg16: if pretrained: weights models.VGG16_Weights.IMAGENET1K_V1 model models.vgg16(weightsweights) in_feat model.classifier[-1].in_features model.classifier[-1] nn.Linear(in_feat, num_classes) elif name resnet50: if pretrained: weights models.ResNet50_Weights.IMAGENET1K_V1 model models.resnet50(weightsweights) in_feat model.fc.in_features model.fc nn.Linear(in_feat, num_classes) elif name vit: if pretrained: weights models.ViT_B_16_Weights.IMAGENET1K_V1 model models.vit_b_16(weightsweights) in_feat model.heads.head.in_features model.heads.head nn.Linear(in_feat, num_classes) else: raise ValueError(fUnsupported model: {name}) return model这段代码的关键是分类头的适配。AlexNet 和 VGG 的分类头在classifier末尾ResNet 在fcViT 在heads.head。如果你用的是其他变体比如 ResNet18、VGG11、ViT-Large只需要修改对应的构造行。预训练权重的选择要注意第一次运行会自动下载权重文件需要网络。如果网络不畅可以先手动下载到~/.cache/torch或指定weights_dir。如果不想用预训练权重把pretrainedFalse即可但训练集太小或训练轮次不足时准确率通常不如微调预训练模型。8. 训练循环与验证train.py是核心脚本包含训练、验证、学习率调整、早停和模型保存。整个逻辑不依赖具体模型结构模型切换后这段代码不需要改动。import json import time import copy import torch import torch.nn as nn from tqdm import tqdm from config import Config from model import build_model from data_utils import build_loaders def validate(model, val_loader, criterion, device): model.eval() running_loss 0.0 correct 0 total 0 with torch.no_grad(): for inputs, labels in tqdm(val_loader, descVal): inputs, labels inputs.to(device), labels.to(device) outputs model(inputs) loss criterion(outputs, labels) running_loss loss.item() * inputs.size(0) _, preds torch.max(outputs, 1) total labels.size(0) correct (preds labels).sum().item() avg_loss running_loss / total acc correct / total return avg_loss, acc def train(cfg): device torch.device(cfg.device) model build_model(cfg.model_name, cfg.num_classes, pretrainedcfg.pretrained) model.to(device) criterion nn.CrossEntropyLoss() optimizer torch.optim.SGD(model.parameters(), lrcfg.lr, momentum0.9, weight_decaycfg.weight_decay) scheduler torch.optim.lr_scheduler.ReduceLROnPlateau( optimizer, modemin, factor0.1, patience3 ) train_loader, val_loader, class_names build_loaders(cfg) best_acc 0.0 best_weights copy.deepcopy(model.state_dict()) patience_counter 0 history [] print(fTraining {cfg.model_name} on {device} ...) for epoch in range(cfg.epochs): model.train() running_loss 0.0 correct 0 total 0 start time.time() for inputs, labels in tqdm(train_loader, descfEpoch {epoch1}/{cfg.epochs}): inputs, labels inputs.to(device), labels.to(device) optimizer.zero_grad() outputs model(inputs) loss criterion(outputs, labels) loss.backward() optimizer.step() running_loss loss.item() * inputs.size(0) _, preds torch.max(outputs, 1) total labels.size(0) correct (preds labels).sum().item() train_loss running_loss / total train_acc correct / total val_loss, val_acc validate(model, val_loader, criterion, device) elapsed time.time() - start scheduler.step(val_loss) lr_now optimizer.param_groups[0][lr] print(fEpoch {epoch1}: train_loss{train_loss:.4f}, train_acc{train_acc:.4f}, fval_loss{val_loss:.4f}, val_acc{val_acc:.4f}, lr{lr_now:.6f}, time{elapsed:.1f}s) history.append({ epoch: epoch 1, train_loss: train_loss, train_acc: train_acc, val_loss: val_loss, val_acc: val_acc, lr: lr_now }) if val_acc best_acc: best_acc val_acc best_weights copy.deepcopy(model.state_dict()) torch.save(best_weights, cfg.save_path) print(fBest model saved to {cfg.save_path}, val_acc{val_acc:.4f}) patience_counter 0 else: patience_counter 1 if patience_counter cfg.patience: print(fEarly stop at epoch {epoch1}) break model.load_state_dict(best_weights) with open(runs/history.json, w, encodingutf-8) as f: json.dump(history, f, indent2) print(fTraining finished. Best val_acc{best_acc:.4f}) if __name__ __main__: cfg Config() train(cfg)这段代码里有几个细节值得说明第一用copy.deepcopy(model.state_dict())保存最佳权重避免后续训练迭代污染最佳结果。第二ReduceLROnPlateau监控验证损失连续三个 epoch 不下降就降低学习率。第三验证过程用torch.no_grad()避免梯度计算。第四训练时开启model.train()验证时切回model.eval()否则 BN 和 Dropout 行为不一致。如果你换用 Adam 优化器可以在config.py里加一个优化器参数if getattr(cfg, optimizer, sgd).lower() adam: optimizer torch.optim.Adam(model.parameters(), lrcfg.lr, weight_decaycfg.weight_decay)9. 各模型训练策略与效果对比把model_name分别改为alexnet、vgg16、resnet50、vit即可重复训练。但四类模型的训练策略有明显差异。AlexNet 参数量约 6000 万结构较浅。在自定义小数据集上训练速度快但容易过拟合建议开启数据增强、调小 batch size、增大 weight_decay。VGG16 参数量约 1.38 亿且前两个全连接层各有 4096 个节点显存占用高。训练时建议使用预训练权重微调全连接层可以换成更小的结构。如果显存不够可以把batch_size降到 8 或 4也可以使用梯度累积。ResNet50 参数量约 2500 万残差连接让训练比较稳定。在中等规模数据集上ResNet 往往比 VGG 更快收敛、效果更好。微调时可以先冻结 backbone只训练最后一层跑几个 epoch 后再解冻所有层。ViT-B/16 参数量约 8600 万但训练策略与 CNN 差异更大。ViT 在大型数据集上表现好在小数据集上直接从头训练容易欠拟合需要更多数据增强、正则化以及更大的训练轮次。微调时学习率通常比 CNN 低建议使用 Lion 或 AdamW并配合 warmup。ViT 对分辨率更敏感如果你的图像细节重要可以尝试 384x384 输入但显存会成倍增长。从训练曲线对比的角度看一个小型数据集上常见现象是AlexNet 和 VGG 收敛较快但最终准确率受限于特征表达ResNet 用深层结构获得更好泛化ViT 在小数据上可能不如 CNN但在数据足够的情况下上限更高。具体数值依赖数据集不要当作绝对结论。10. 模型推理与导出部署训练完成后checkpoint/best_model.pth保存的是state_dict不是完整模型。推理时要先构建模型结构再加载权重。写一个推理脚本import torch from PIL import Image from torchvision import transforms from model import build_model from config import Config cfg Config() device torch.device(cfg.device if torch.device(cfg.device).type ! cpu else cpu) model build_model(cfg.model_name, cfg.num_classes, pretrainedFalse) model.load_state_dict(torch.load(cfg.save_path, map_locationdevice)) model.to(device).eval() transform transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(cfg.input_size), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) def predict(image_path, class_names): img Image.open(image_path).convert(RGB) tensor transform(img).unsqueeze(0).to(device) with torch.no_grad(): outputs model(tensor) probs torch.softmax(outputs, dim1) conf, idx torch.max(probs, dim1) return class_names[idx.item()], conf.item() if __name__ __main__: # 替换为自己的图片路径 image_path data/val/class0/test.jpg # 类别列表要在训练时从 ImageFolder 获取并保存 class_names [class0, class1, class2] cls, conf predict(image_path, class_names) print(fPredicted: {cls}, confidence: {conf:.4f})推理时有个常见坑加载权重时如果cfg.pretrainedTrue训练脚本用预训练权重初始化的模型结构而推理脚本设置pretrainedFalse只是不下载预训练权重结构本身一致所以能正常加载。但必须保持num_classes一致。导出 ONNX 用于加速推理代码import torch from config import Config from model import build_model cfg Config() model build_model(cfg.model_name, cfg.num_classes, pretrainedFalse) model.load_state_dict(torch.load(cfg.save_path, map_locationcpu)) model.eval() dummy torch.randn(1, 3, cfg.input_size, cfg.input_size) torch.onnx.export( model, dummy, cfg.onnx_path, opset_version12, input_names[input], output_names[output] ) print(ONNX exported to, cfg.onnx_path)用 ONNXRuntime 推理时注意输入张量的归一化要与训练时完全一致否则准确率会下降。ONNX 导出后还可以用onnxruntime或 TensorRT 部署演示样例如下import numpy as np import onnxruntime as ort from PIL import Image from torchvision import transforms sess ort.InferenceSession(checkpoint/best_model.onnx) img Image.open(data/val/class0/test.jpg).convert(RGB) transform transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) tensor transform(img).unsqueeze(0).numpy() name sess.get_inputs()[0].name probs sess.run(None, {name: tensor})[0][0] cls_idx np.argmax(probs) conf probs[cls_idx] print(ONNX predicted class index:, cls_idx, confidence:, conf)ONNX 模型可以在 CPU 上获得稳定的推理速度是入门部署的一个良好选择。11. 接口 API 与批量任务训练好的模型可以通过 Flask 提供一个轻量的 HTTP 接口。app.py示例import io import torch from flask import Flask, request, jsonify from PIL import Image from torchvision import transforms from model import build_model from config import Config cfg Config() device torch.device(cuda if torch.cuda.is_available() else cpu) model build_model(cfg.model_name, cfg.num_classes, pretrainedFalse) model.load_state_dict(torch.load(cfg.save_path, map_locationdevice)) model.to(device).eval() transform transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(cfg.input_size), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) app Flask(__name__) app.route(/predict, methods[POST]) def predict(): if image not in request.files: return jsonify({error: no image}), 400 file request.files[image] img Image.open(io.BytesIO(file.read())).convert(RGB) tensor transform(img).unsqueeze(0).to(device) with torch.no_grad(): outputs model(tensor) probs torch.softmax(outputs, dim1) conf, idx torch.max(probs, dim1) # 这里假定类别列表为空实际项目应传入类别名 return jsonify({ class_index: idx.item(), confidence: round(conf.item(), 4), class_name: unknown }) if __name__ __main__: app.run(hostcfg.api_host, portcfg.api_port)启动 APIpython app.py用 curl 测试curl -X POST http://127.0.0.1:5000/predict \ -F imagedata/val/class0/test.jpg返回结果类似{ class_index: 0, confidence: 0.9821, class_name: unknown }批量任务不需要改服务端代码。客户端可以循环读取目录里的图片逐张 POST 请求然后汇总结果import os import requests url http://127.0.0.1:5000/predict image_dir batch_imgs results [] for fname in os.listdir(image_dir): path os.path.join(image_dir, fname) with open(path, rb) as f: response requests.post(url, files{image: (fname, f, image/jpeg)}, timeout30) if response.status_code 200: results.append({file: fname, **response.json()}) else: results.append({file: fname, error: response.text}) for r in results: print(r)批量任务要注意两点第一单张请求能跑通不代表并发流畅Flask 默认开发服务器不支持高并发生产环境至少换成 gunicorn gevent第二如果需要更高的吞吐量建议在客户端做多线程并发或者服务端把图片转为 base64 放在 JSON 里传输避免 multipart 解析开销。12. 资源占用与性能观察训练过程中的显存占用可以通过nvidia-smi观察watch -n 1 nvidia-smi实际显存占用由几个因素决定模型参数、优化器状态、中间激活值和 batch size。同一个模型在 224x224 输入下batch_size16的 ResNet50 通常需要 4GB 左右显存VGG16 会明显更高ViT 视 patch 和深度而定。这里的数值只是经验范围你的环境里实际多少需要自己观察。降低显存占用的常见手段包括降低 batch size使用小分辨率输入使用混合精度训练PyTorch 的torch.cuda.amp开启梯度累积使用torchvision里更轻量的变体比如resnet18或vit_b_16。性能观察不能只看显存。训练耗时同样受 CPU 解码速度影响。num_workers设置过小会导致 GPU 空转设置过大会增加内存占用。可以先从num_workers2开始测试再逐步提高。如果数据量不大pin_memoryTrue可以减少主机到显存传输的阻塞时间。推理性能方面ONNX Runtime 在 CPU 上的速度通常优于原始 PyTorch 的 eager 模式特别是固定批量时。你可以用以下命令快速比较两种方式的耗时python -c import torch; import onnxruntime; print(compare by timeit script)不做过度优化先把流程跑通再针对瓶颈做剪枝或量化。13. 常见问题与排查方法问题现象可能原因排查方式解决方案torch.cuda.is_available()返回 FalsePyTorch 安装的是 CPU 版显卡驱动或 CUDA 版本不匹配运行python -c import torch; print(torch.__version__)检查后端从 PyTorch 官网重新安装匹配 CUDA 的版本更新显卡驱动数据加载时报EOFError或文件夹不存在数据目录结构不符合ImageFolder要求确认train_dir中存在子文件夹且子文件夹内是图片整理目录结构确保每个类别一个文件夹训练时显存溢出CUDA out of memorybatch size 过大或输入分辨率过高观察nvidia-smi的显存占用降低 batch size或使用torch.utils.data.dataloader的pin_memory配置尝试混合精度模型加载权重时报size mismatch训练时的num_classes与当前模型的num_classes不一致打印模型最后分类层的维度检查Config.num_classes和类别数是否一致检查分类头替换是否正确预训练权重下载失败网络问题或证书问题在第一次运行前手动下载权重文件设置torch.hub.set_dir(./weights)并手动放置权重或关闭预训练API 请求返回 400请求没有带image文件字段用 Postman 或用 curl-F参数测试检查客户端上传字段名与 Flask 代码一致批量任务中途卡住某张图片损坏或超时增加日志打印当前文件加异常捕获跳过异常图片设置requests超时使用队列任务管理训练准确率很低数据预处理不一致学习率不合适模型欠拟合先跑 5 个 epoch 观察 loss 是否下降检查归一化参数调低学习率或换优化器增加训练轮次问题排查的原则是先看日志再看显存最后才怀疑模型结构。大部分训练失败都出在数据路径、归一化参数和分类头维度这三个地方。14. 最佳实践与使用建议第一次跑模板时不要直接用完整数据集。建议先复制一个小型子集比如每个类别 20 张图跑 3 个 epoch验证从数据读取到梯度更新再到模型保存全链路没有问题。这样做可以快速暴露配置错误避免浪费几个小时才发现路径写错。训练过程中保留一份history.json。这个文件里记录了每个 epoch 的损失和准确率方便画训练曲线、调参与写报告。后期可以加一行代码用 matplotlib 输出曲线不过不是必需功能。模型、日志、输入数据、输出结果建议分目录管理。比如checkpoint/放权重runs/放日志data/放图片output/放批处理结果。模板已经按这个思路设计。批量任务必须加日志和失败重试。一个最简单的重试策略是循环发送请求时记录失败文件结束后统一重试两次。如果图片数量很多考虑用 Redis 或 RabbitMQ 搭队列但那是后话。接口服务如果只在本地使用监听127.0.0.1就够了。如果要在局域网内访问再改成0.0.0.0但必须加上简单的 token 校验避免别人随意调用。加 token 的做法是请求头带一个固定值服务端校验。涉及人脸、声音、版权素材时确认授权是第一优先级。不要在公开数据集上训练后直接商用除非数据集和模型权重许可证允许。发布演示 Demo 时对图片做脱敏。15. 总结与下一步这份模板最值得尝试的点是用同一个train.py把 AlexNet、VGG、ResNet、ViT 四个模型全部串起来它能帮你快速理解不同 CNN 和 Transformer 架构在代码层面上的差异。最先该验证的功能不是最终准确率而是“改一个model_name模型结构是否成功切换并开始训练”。最容易踩的坑是分类头替换时维度不匹配以及数据目录结构不符合ImageFolder约定。后续扩展方向很明确把build_model里加入更多模型如 ResNet18、MobileNetV3、EfficientNet、Swin Transformer训练循环里加入混合精度和分布式部署端用 FastAPI 替换 Flask结合 ONNX Runtime 或 TensorRT 做生产级服务。你甚至可以把这份模板改造成图像特征提取器用于检索或对比学习。先把基础流程跑通再按需求往工程化方向演进。
网站建设高端定制企业官网
RELATED

相关资讯

更多精彩内容,欢迎继续阅读

较早相关资讯

最新相关资讯

【OpenClaw从入门到精通】第87篇:用 TaoToken 统一 Key 跑通你的第一个自定义 Agent:YAML 配置到 Python 运行完整实战 2026/9/26 12:20:52

【OpenClaw从入门到精通】第87篇:用 TaoToken 统一 Key 跑通你的第一个自定义 Agent:YAML 配置到 Python 运行完整实战

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

阅读更多 →
异步RL架构全拆解:三台可扩容机器如何重构Agent训练循环 2026/9/26 12:20:32

异步RL架构全拆解:三台可扩容机器如何重构Agent训练循环

最近这套“小米 MiMo-V2.6 全异步 RL”架构在 Agent 训练相关的讨论里反复出现,标题里的“每小时 3 万美元”先抓眼球,但真正值得研究的,是它把传统 RL 里那种“跑一步等一步”的 Agent 训练循环,改成了三组可以单独水平扩容的机器…

阅读更多 →
Epoll模型详解:从epoll_ctl到epoll_wait的Linux IO多路复用实践 2026/9/26 12:20:32

Epoll模型详解:从epoll_ctl到epoll_wait的Linux IO多路复用实践

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

阅读更多 →
网页模板HTML源码下载与改造:免费源码选型、避坑与上线全攻略 2026/9/26 12:20:32

网页模板HTML源码下载与改造:免费源码选型、避坑与上线全攻略

简介:这是一套面向网页开发初学者的基础HTML模板源码,由样式表、结构文档、交互脚本及图片资源共同构成,适合用于快速搭建静态网站或学习HTML/CSS/JavaScript协作流程。压缩包共9个文件,包括template.html、styles.css、script.js…

阅读更多 →
解决 Unit nginx.service not found:systemd 服务排查与 Nginx 启动修复指南 2026/9/26 12:20:26

解决 Unit nginx.service not found:systemd 服务排查与 Nginx 启动修复指南

如果你在服务器上敲下 systemctl restart nginx ,迎面却撞见一行红字: Failed to restart nginx.service: Unit nginx.service not found ,别急着怀疑人生。这个报错十有八九不是 nginx 本身坏了,而是 systemd 根本找不到名为…

阅读更多 →
Cell封面|衰老生物学开源AI工具包 2026/9/26 12:20:26

Cell封面|衰老生物学开源AI工具包

英矽智能在Cell研究中发布开源长寿AI工具包#长寿 #衰老 #基准 #大语言模型 #微调 #衰老时钟 #多组学 #衰老保护剂 #基础模型 #靶点发现 Source: Unite.AI英矽智能于2026年9月17日宣布,在Cell期刊发表研究,推出面向衰老生物学的开源AI工具包:…

阅读更多 →

今日资讯

本周资讯

本月资讯

看完文章仍有疑问?

联系尧图顾问,获取一对一建站咨询

立即免费咨询 📞 400-888-8888
📞 ✉