廢棄物分類數(shù)據(jù)集實(shí)戰(zhàn):4,800張圖從訓(xùn)練到部署)
簡(jiǎn)介本資源為面向計(jì)算機(jī)視覺初學(xué)者與圖像分類實(shí)踐者的真實(shí)廢棄物圖像分類數(shù)據(jù)集覆蓋紙板、食品有機(jī)物、玻璃、金屬、雜項(xiàng)垃圾、紙張、塑料、紡織品垃圾和植被共9個(gè)類別適合用于分類網(wǎng)絡(luò)訓(xùn)練、遷移學(xué)習(xí)驗(yàn)證及垃圾分類相關(guān)課程設(shè)計(jì)。數(shù)據(jù)已完成預(yù)處理可直接作為分類網(wǎng)絡(luò)輸入并已劃分訓(xùn)練集與測(cè)試集各類別圖片分目錄存放便于快速構(gòu)建實(shí)驗(yàn)流程。壓縮包共2000個(gè)文件以1998張jpg圖像為主體另含1個(gè)json標(biāo)注文件與1個(gè)Python可視化腳本整體約155.99MB運(yùn)行show腳本即可直觀查看樣本分布與類別效果。目前已有65人學(xué)習(xí)下載讀者可借此省去數(shù)據(jù)采集與清洗成本將精力集中于模型結(jié)構(gòu)改進(jìn)、參數(shù)調(diào)優(yōu)與結(jié)果對(duì)比同時(shí)結(jié)合標(biāo)注文件理解類別映射關(guān)系為后續(xù)分割或分類任務(wù)提供可靠的數(shù)據(jù)基礎(chǔ)。1. 真實(shí)廢棄物分類數(shù)據(jù)集4,800 張標(biāo)注圖能跑出什么名堂拿到「生活中真實(shí)廢棄物圖像分類數(shù)據(jù)集」這個(gè)標(biāo)題多數(shù)人第一反應(yīng)是去找下載鏈接但真正決定項(xiàng)目成敗的是標(biāo)注質(zhì)量與類別體系能不能對(duì)上你要落地的場(chǎng)景。這個(gè)數(shù)據(jù)集約 4,800 張、已完成標(biāo)注屬于中小規(guī)模圖像分類數(shù)據(jù)集適合做垃圾分類識(shí)別、智能回收箱、環(huán)衛(wèi)巡檢等方向的模型驗(yàn)證與原型開發(fā)。它解決的核心問題是讓你跳過最耗時(shí)的采集與標(biāo)注環(huán)節(jié)直接進(jìn)入模型訓(xùn)練和效果調(diào)優(yōu)。適合誰想快速驗(yàn)證圖像分類算法的新手、需要 baseline 做對(duì)比的算法工程師、以及做課程設(shè)計(jì)或產(chǎn)品 demo 的開發(fā)者。但別指望它直接產(chǎn)出生產(chǎn)級(jí)模型——4,800 張的體量決定了它更適合跑通流程、驗(yàn)證思路而不是追求 SOTA 精度。2. 廢棄物圖像分類的數(shù)據(jù)集拆解與模型選型2.1 先搞清楚 4,800 張圖里到底有什么在動(dòng)手寫任何訓(xùn)練代碼之前必須先把數(shù)據(jù)集的結(jié)構(gòu)摸清楚。真實(shí)廢棄物圖像分類數(shù)據(jù)集通常按類別分文件夾存放目錄結(jié)構(gòu)類似dataset/train/plastic/、dataset/train/paper/這種形式。你需要確認(rèn)三件事類別數(shù)量、每類樣本數(shù)、圖像尺寸分布。import os from collections import Counter from PIL import Image data_dir dataset/train class_counts {} size_dist Counter() for cls in sorted(os.listdir(data_dir)): cls_path os.path.join(data_dir, cls) if not os.path.isdir(cls_path): continue imgs [f for f in os.listdir(cls_path) if f.lower().endswith((.jpg, .png, .jpeg))] class_counts[cls] len(imgs) for img_name in imgs[:20]: # 每類抽樣20張看尺寸 with Image.open(os.path.join(cls_path, img_name)) as im: size_dist[im.size] 1 print(類別分布, class_counts) print(尺寸分布抽樣, size_dist.most_common(5)) total sum(class_counts.values()) print(f總圖片數(shù){total}) for cls, cnt in class_counts.items(): print(f {cls}: {cnt} ({cnt/total*100:.1f}%))這段腳本做的是類別均衡性檢查。參數(shù)說明data_dir指向訓(xùn)練集根目錄腳本會(huì)自動(dòng)遍歷子文件夾。重點(diǎn)看輸出里的百分比——如果某個(gè)類別占比超過 40% 或低于 5%就存在類別不平衡問題后續(xù)訓(xùn)練需要加WeightedRandomSampler或做數(shù)據(jù)增強(qiáng)補(bǔ)償。尺寸分布則決定你統(tǒng)一 resize 到多少常見做法是 224×224 或 256×256。2.2 模型選型從 ResNet 到 Transformer 的取舍2024 年之后圖像分類模型的選擇面比幾年前寬了很多。對(duì)于 4,800 張這個(gè)量級(jí)我的建議是分三檔模型參數(shù)量適用場(chǎng)景預(yù)期準(zhǔn)確率5折ResNet-1811M快速 baselineCPU 也能跑82%88%EfficientNet-B05.3M精度與速度平衡85%91%ViT-B/16預(yù)訓(xùn)練86M數(shù)據(jù)增強(qiáng)充分時(shí)88%93%Swin-Tiny28M想要 Transformer 但顯存有限87%92%選型的核心邏輯不是「哪個(gè)最新」而是「你的數(shù)據(jù)量和算力撐得住哪個(gè)」。4,800 張圖直接訓(xùn) ViT 從零開始基本會(huì)過擬合但如果用 ImageNet 預(yù)訓(xùn)練權(quán)重做遷移學(xué)習(xí)ViT 系列反而可能比 CNN 高 23 個(gè)點(diǎn)。我一般會(huì)先用 ResNet-18 跑一個(gè) baseline確認(rèn)數(shù)據(jù) pipeline 沒問題再換 EfficientNet 或 Swin 做精度提升。import torch import torchvision.models as models import torch.nn as nn def build_model(num_classes, archresnet18, pretrainedTrue): if arch resnet18: model models.resnet18(weightsmodels.ResNet18_Weights.DEFAULT if pretrained else None) model.fc nn.Linear(model.fc.in_features, num_classes) elif arch efficientnet_b0: model models.efficientnet_b0(weightsmodels.EfficientNet_B0_Weights.DEFAULT if pretrained else None) model.classifier[1] nn.Linear(model.classifier[1].in_features, num_classes) elif arch swin_t: model models.swin_t(weightsmodels.Swin_T_Weights.DEFAULT if pretrained else None) model.head nn.Linear(model.head.in_features, num_classes) else: raise ValueError(f不支持的架構(gòu): {arch}) return model # 假設(shè)有6個(gè)廢棄物類別 model build_model(num_classes6, archefficientnet_b0) print(f可訓(xùn)練參數(shù){sum(p.numel() for p in model.parameters() if p.requires_grad):,})關(guān)鍵點(diǎn)在于替換分類頭ResNet 換fcEfficientNet 換classifier[1]Swin 換head。pretrainedTrue時(shí)加載 ImageNet 權(quán)重這是小數(shù)據(jù)集能訓(xùn)出可用模型的前提。參數(shù)量打印出來是為了確認(rèn)你改對(duì)了層——如果可訓(xùn)練參數(shù)接近全量參數(shù)說明預(yù)訓(xùn)練權(quán)重沒加載成功。2.3 數(shù)據(jù)增強(qiáng)策略別讓 4,800 張圖白瞎了4,800 張圖做分類數(shù)據(jù)增強(qiáng)不是可選項(xiàng)而是必選項(xiàng)。但廢棄物圖像有個(gè)特殊性翻轉(zhuǎn)和旋轉(zhuǎn)通常安全顏色抖動(dòng)要謹(jǐn)慎——塑料瓶被調(diào)成奇怪色調(diào)可能影響模型對(duì)材質(zhì)的判斷。from torchvision import transforms train_transform transforms.Compose([ transforms.RandomResizedCrop(224, scale(0.7, 1.0)), transforms.RandomHorizontalFlip(p0.5), transforms.RandomVerticalFlip(p0.2), transforms.RandomRotation(15), transforms.ColorJitter(brightness0.2, contrast0.2, saturation0.1, hue0.05), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), transforms.RandomErasing(p0.1, scale(0.02, 0.1)), ]) val_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]), ])參數(shù)說明RandomResizedCrop的scale(0.7, 1.0)表示隨機(jī)裁剪原圖 70%100% 的區(qū)域再縮放模擬不同拍攝距離。ColorJitter的hue0.05控制得很小就是為了避免顏色失真。RandomErasing模擬遮擋場(chǎng)景對(duì)廢棄物識(shí)別很有用——實(shí)際場(chǎng)景中垃圾經(jīng)常被部分遮擋。驗(yàn)證集只用 resize centercrop保證評(píng)估一致性。3. 從零跑通訓(xùn)練命令行、參數(shù)與監(jiān)控3.1 訓(xùn)練腳本的核心結(jié)構(gòu)一個(gè)能復(fù)現(xiàn)的訓(xùn)練腳本需要包含數(shù)據(jù)加載、模型構(gòu)建、損失函數(shù)、優(yōu)化器、學(xué)習(xí)率調(diào)度、訓(xùn)練循環(huán)、驗(yàn)證循環(huán)、模型保存。下面是一個(gè)精簡(jiǎn)但完整的版本。import torch from torch.utils.data import DataLoader from torchvision.datasets import ImageFolder from torch.optim import AdamW from torch.optim.lr_scheduler import CosineAnnealingLR import time def train_one_epoch(model, loader, criterion, optimizer, device): model.train() running_loss, correct, total 0.0, 0, 0 for imgs, labels in loader: imgs, labels imgs.to(device), labels.to(device) optimizer.zero_grad() outputs model(imgs) loss criterion(outputs, labels) loss.backward() optimizer.step() running_loss loss.item() * imgs.size(0) correct (outputs.argmax(1) labels).sum().item() total imgs.size(0) return running_loss / total, correct / total torch.no_grad() def evaluate(model, loader, criterion, device): model.eval() running_loss, correct, total 0.0, 0, 0 for imgs, labels in loader: imgs, labels imgs.to(device), labels.to(device) outputs model(imgs) loss criterion(outputs, labels) running_loss loss.item() * imgs.size(0) correct (outputs.argmax(1) labels).sum().item() total imgs.size(0) return running_loss / total, correct / total # 主流程 device torch.device(cuda if torch.cuda.is_available() else cpu) train_ds ImageFolder(dataset/train, transformtrain_transform) val_ds ImageFolder(dataset/val, transformval_transform) train_loader DataLoader(train_ds, batch_size32, shuffleTrue, num_workers4, pin_memoryTrue) val_loader DataLoader(val_ds, batch_size32, shuffleFalse, num_workers4, pin_memoryTrue) num_classes len(train_ds.classes) model build_model(num_classes, archefficientnet_b0).to(device) criterion torch.nn.CrossEntropyLoss(label_smoothing0.1) optimizer AdamW(model.parameters(), lr3e-4, weight_decay1e-4) scheduler CosineAnnealingLR(optimizer, T_max30) best_acc 0.0 for epoch in range(30): t0 time.time() tr_loss, tr_acc train_one_epoch(model, train_loader, criterion, optimizer, device) val_loss, val_acc evaluate(model, val_loader, criterion, device) scheduler.step() if val_acc best_acc: best_acc val_acc torch.save(model.state_dict(), best_model.pth) print(fEpoch {epoch1:02d} | train_loss{tr_loss:.4f} acc{tr_acc:.3f} | fval_loss{val_loss:.4f} acc{val_acc:.3f} | {time.time()-t0:.1f}s) print(f最佳驗(yàn)證準(zhǔn)確率{best_acc:.4f})邏輯說明label_smoothing0.1緩解過擬合對(duì)小數(shù)據(jù)集效果明顯。AdamW的lr3e-4是遷移學(xué)習(xí)的常用起點(diǎn)weight_decay1e-4做正則化。CosineAnnealingLR讓學(xué)習(xí)率按余弦曲線下降T_max30對(duì)應(yīng)總 epoch 數(shù)。每個(gè) epoch 結(jié)束后比較驗(yàn)證準(zhǔn)確率只保存最好的模型——這是防止過擬合的后悔藥。3.2 訓(xùn)練過程中的關(guān)鍵監(jiān)控指標(biāo)光看 loss 和 accuracy 不夠你還需要關(guān)注訓(xùn)練/驗(yàn)證 loss 的差距判斷過擬合、每類準(zhǔn)確率判斷類別不平衡影響、學(xué)習(xí)率變化確認(rèn)調(diào)度器生效。建議加一個(gè)簡(jiǎn)單的混淆矩陣輸出。from sklearn.metrics import classification_report, confusion_matrix torch.no_grad() def detailed_eval(model, loader, device, class_names): model.eval() all_preds, all_labels [], [] for imgs, labels in loader: imgs imgs.to(device) preds model(imgs).argmax(1).cpu().numpy() all_preds.extend(preds) all_labels.extend(labels.numpy()) print(classification_report(all_labels, all_preds, target_namesclass_names)) print(混淆矩陣) print(confusion_matrix(all_labels, all_preds)) detailed_eval(model, val_loader, device, train_ds.classes)classification_report會(huì)輸出每類的 precision、recall、f1-score。如果某個(gè)類別 recall 特別低說明模型對(duì)這個(gè)類別識(shí)別能力差可能需要補(bǔ)充該類樣本或調(diào)整采樣權(quán)重?;煜仃噭t告訴你哪些類別容易被搞混——比如「紙類」和「紙板」如果混淆嚴(yán)重說明類別定義本身可能有問題。3.3 用 TensorBoard 或 wandb 記錄實(shí)驗(yàn)命令行打印只能看當(dāng)前狀態(tài)做對(duì)比實(shí)驗(yàn)需要可視化工具。最輕量的方案是 TensorBoard。from torch.utils.tensorboard import SummaryWriter writer SummaryWriter(runs/waste_classify_exp1) for epoch in range(30): tr_loss, tr_acc train_one_epoch(model, train_loader, criterion, optimizer, device) val_loss, val_acc evaluate(model, val_loader, criterion, device) writer.add_scalars(Loss, {train: tr_loss, val: val_loss}, epoch) writer.add_scalars(Accuracy, {train: tr_acc, val: val_acc}, epoch) writer.add_scalar(LR, optimizer.param_groups[0][lr], epoch) scheduler.step() writer.close()啟動(dòng)命令tensorboard --logdirruns。參數(shù)說明add_scalars把訓(xùn)練和驗(yàn)證曲線畫在同一張圖上方便對(duì)比。add_scalar(LR, ...)確認(rèn)學(xué)習(xí)率按預(yù)期下降。如果 val loss 在某個(gè) epoch 后開始上升而 train loss 繼續(xù)下降就是過擬合信號(hào)需要提前停止或加強(qiáng)正則化。4. 避坑與排查廢棄物分類訓(xùn)練中最容易翻車的 5 個(gè)點(diǎn)4.1 類別文件夾命名混亂導(dǎo)致標(biāo)簽錯(cuò)位現(xiàn)象訓(xùn)練時(shí)準(zhǔn)確率始終在隨機(jī)水平附近比如 6 分類一直在 16% 左右loss 不下降。原因ImageFolder按文件夾名排序生成標(biāo)簽如果文件夾命名有中文、空格或大小寫不一致可能導(dǎo)致標(biāo)簽映射混亂。更隱蔽的情況是某些文件夾里混入了其他類別的圖片。解決訓(xùn)練前先打印train_ds.classes和train_ds.class_to_idx確認(rèn)類別列表和映射關(guān)系符合預(yù)期。再用腳本檢查每個(gè)文件夾內(nèi)是否有異常圖片比如尺寸為 0 或無法打開的文件。for cls in train_ds.classes: cls_path os.path.join(dataset/train, cls) for f in os.listdir(cls_path): fp os.path.join(cls_path, f) try: with Image.open(fp) as im: im.verify() except Exception as e: print(f損壞文件: {fp} - {e})4.2 驗(yàn)證集與訓(xùn)練集分布不一致現(xiàn)象驗(yàn)證準(zhǔn)確率遠(yuǎn)低于訓(xùn)練準(zhǔn)確率且差距持續(xù)擴(kuò)大。原因劃分驗(yàn)證集時(shí)沒有做分層采樣導(dǎo)致某些類別在驗(yàn)證集中占比過高或過低?;蛘唑?yàn)證集的圖片來自不同拍攝條件比如訓(xùn)練集是白底圖驗(yàn)證集是實(shí)景圖。解決用sklearn.model_selection.train_test_split的stratify參數(shù)做分層劃分保證每個(gè)類別在訓(xùn)練集和驗(yàn)證集中的比例一致。如果數(shù)據(jù)集本身已經(jīng)分好 train/val先檢查兩邊的類別分布是否接近。from sklearn.model_selection import train_test_split import numpy as np all_files, all_labels [], [] for idx, cls in enumerate(sorted(os.listdir(dataset/train))): cls_path os.path.join(dataset/train, cls) for f in os.listdir(cls_path): all_files.append(os.path.join(cls_path, f)) all_labels.append(idx) X_train, X_val, y_train, y_val train_test_split( all_files, all_labels, test_size0.2, stratifyall_labels, random_state42 ) print(訓(xùn)練集類別分布, np.bincount(y_train)) print(驗(yàn)證集類別分布, np.bincount(y_val))4.3 圖像尺寸不統(tǒng)一導(dǎo)致 DataLoader 報(bào)錯(cuò)現(xiàn)象訓(xùn)練時(shí)報(bào)RuntimeError: stack expects each tensor to be equal size。原因數(shù)據(jù)集中存在尺寸差異極大的圖片而 transform 中沒有做統(tǒng)一 resize或者某些圖片是灰度圖/ RGBA 四通道圖。解決在 transform 中強(qiáng)制 resize 到固定尺寸并在數(shù)據(jù)集類中加一個(gè)轉(zhuǎn)換步驟把灰度圖和 RGBA 圖統(tǒng)一轉(zhuǎn)成 RGB。from PIL import Image class WasteDataset(torch.utils.data.Dataset): def __init__(self, file_list, labels, transformNone): self.file_list file_list self.labels labels self.transform transform def __len__(self): return len(self.file_list) def __getitem__(self, idx): img Image.open(self.file_list[idx]).convert(RGB) # 關(guān)鍵統(tǒng)一轉(zhuǎn)RGB label self.labels[idx] if self.transform: img self.transform(img) return img, label4.4 學(xué)習(xí)率設(shè)太大導(dǎo)致 loss 震蕩不收斂現(xiàn)象訓(xùn)練初期 loss 劇烈震蕩甚至出現(xiàn) NaN。原因遷移學(xué)習(xí)時(shí)用了從零訓(xùn)練的學(xué)習(xí)率比如 0.1而預(yù)訓(xùn)練模型需要更小的學(xué)習(xí)率來微調(diào)。解決遷移學(xué)習(xí)場(chǎng)景下AdamW 的 lr 建議從 1e-4 到 3e-4 開始試。如果 loss 仍然震蕩降到 1e-5。另外可以加梯度裁剪。torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0)這行代碼放在loss.backward()之后、optimizer.step()之前防止梯度爆炸。4.5 顯存不夠?qū)е掠?xùn)練中斷現(xiàn)象CUDA out of memory。原因batch_size 太大或者模型參數(shù)量超出顯存容量。解決優(yōu)先降 batch_size從 32 降到 16 或 8其次考慮混合精度訓(xùn)練。from torch.cuda.amp import autocast, GradScaler scaler GradScaler() for imgs, labels in train_loader: imgs, labels imgs.to(device), labels.to(device) optimizer.zero_grad() with autocast(): outputs model(imgs) loss criterion(outputs, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()混合精度訓(xùn)練能省 30%50% 顯存對(duì) 4,800 張圖這個(gè)量級(jí)完全夠用。注意autocast只包前向傳播反向傳播用scaler處理。5. 把 4,800 張圖用到極致進(jìn)階技巧與驗(yàn)證方法5.1 用交叉驗(yàn)證榨干小數(shù)據(jù)集的每一張圖4,800 張圖做單次劃分驗(yàn)證集可能只有 8001,000 張?jiān)u估結(jié)果波動(dòng)大。5 折交叉驗(yàn)證能讓每張圖都參與驗(yàn)證一次得到更穩(wěn)定的性能估計(jì)。from sklearn.model_selection import StratifiedKFold skf StratifiedKFold(n_splits5, shuffleTrue, random_state42) fold_results [] for fold, (train_idx, val_idx) in enumerate(skf.split(all_files, all_labels)): train_files [all_files[i] for i in train_idx] val_files [all_files[i] for i in val_idx] train_labels [all_labels[i] for i in train_idx] val_labels [all_labels[i] for i in val_idx] train_ds WasteDataset(train_files, train_labels, transformtrain_transform) val_ds WasteDataset(val_files, val_labels, transformval_transform) train_loader DataLoader(train_ds, batch_size32, shuffleTrue, num_workers4) val_loader DataLoader(val_ds, batch_size32, shuffleFalse, num_workers4) model build_model(num_classes, archefficientnet_b0).to(device) optimizer AdamW(model.parameters(), lr3e-4, weight_decay1e-4) scheduler CosineAnnealingLR(optimizer, T_max20) best_acc 0.0 for epoch in range(20): train_one_epoch(model, train_loader, criterion, optimizer, device) _, val_acc evaluate(model, val_loader, criterion, device) scheduler.step() best_acc max(best_acc, val_acc) fold_results.append(best_acc) print(fFold {fold1}: best_acc{best_acc:.4f}) print(f5折平均準(zhǔn)確率{np.mean(fold_results):.4f} ± {np.std(fold_results):.4f})這個(gè)腳本會(huì)跑 5 次完整訓(xùn)練每次用不同的驗(yàn)證集劃分。最終報(bào)告的是平均準(zhǔn)確率和標(biāo)準(zhǔn)差。標(biāo)準(zhǔn)差如果超過 3 個(gè)點(diǎn)說明模型對(duì)數(shù)據(jù)劃分敏感需要檢查數(shù)據(jù)分布或增加正則化。5.2 用混淆矩陣定位「重災(zāi)區(qū)」類別交叉驗(yàn)證跑完后挑一個(gè) fold 的模型輸出混淆矩陣重點(diǎn)看哪些類別被系統(tǒng)性搞混。import matplotlib.pyplot as plt import seaborn as sns from sklearn.metrics import confusion_matrix cm confusion_matrix(all_labels_fold, all_preds_fold) plt.figure(figsize(8, 6)) sns.heatmap(cm, annotTrue, fmtd, cmapBlues, xticklabelstrain_ds.classes, yticklabelstrain_ds.classes) plt.xlabel(預(yù)測(cè)) plt.ylabel(真實(shí)) plt.title(廢棄物分類混淆矩陣) plt.tight_layout() plt.savefig(confusion_matrix.png, dpi150)如果發(fā)現(xiàn)「塑料瓶」和「玻璃瓶」混淆嚴(yán)重說明模型對(duì)材質(zhì)紋理的區(qū)分能力不足。解決辦法針對(duì)這兩個(gè)類別補(bǔ)充更多樣本或者在數(shù)據(jù)增強(qiáng)中加強(qiáng)紋理相關(guān)的變換比如RandomGrayscale讓模型更關(guān)注形狀而非顏色。5.3 導(dǎo)出 ONNX 做推理驗(yàn)證訓(xùn)練完的模型最終要部署導(dǎo)出 ONNX 是驗(yàn)證模型可用性的關(guān)鍵一步。import torch.onnx model.eval() dummy_input torch.randn(1, 3, 224, 224).to(device) torch.onnx.export( model, dummy_input, waste_classify.onnx, input_names[input], output_names[output], dynamic_axes{input: {0: batch_size}, output: {0: batch_size}}, opset_version13 ) print(ONNX 導(dǎo)出完成) # 驗(yàn)證 ONNX 推理結(jié)果與 PyTorch 一致 import onnxruntime as ort import numpy as np sess ort.InferenceSession(waste_classify.onnx) test_img torch.randn(1, 3, 224, 224).numpy() onnx_out sess.run(None, {input: test_img})[0] torch_out model(torch.from_numpy(test_img).to(device)).cpu().detach().numpy() print(f最大差異{np.abs(onnx_out - torch_out).max():.6f})dynamic_axes讓導(dǎo)出的模型支持可變 batch size部署時(shí)更靈活。最后對(duì)比 ONNX 和 PyTorch 的輸出差異正常應(yīng)該在 1e-5 以內(nèi)。如果差異過大檢查opset_version是否兼容你的推理環(huán)境。5.4 一個(gè)我踩過的坑別在驗(yàn)證集上調(diào)超參數(shù)早期做這個(gè)數(shù)據(jù)集時(shí)我習(xí)慣在驗(yàn)證集上試不同的學(xué)習(xí)率和 batch_size選驗(yàn)證集準(zhǔn)確率最高的那組。結(jié)果模型在測(cè)試集上表現(xiàn)遠(yuǎn)低于預(yù)期——因?yàn)轵?yàn)證集被「偷看」了太多次已經(jīng)失去了評(píng)估的客觀性。正確做法是從訓(xùn)練集里再切一小塊做驗(yàn)證集用于調(diào)參真正的測(cè)試集只在最后評(píng)估一次。如果數(shù)據(jù)量實(shí)在不夠至少保證調(diào)參時(shí)看的是交叉驗(yàn)證的平均結(jié)果而不是單次劃分的驗(yàn)證集。希望幫到你。本文還有配套的精品資源點(diǎn)擊獲取