feat: 添加模型转换工具(Mild v13 PyTorch->bin导出)
- tools/export_mild_v13.py: PyTorch .pth -> 自定义 .bin 导出脚本 - tools/inspect_pth.py: 检查 .pth 权重文件结构 - tools/best_v13_mild_r2.pth: 训练好的 PyTorch 权重 - .gitignore: 追加 build 变体目录和生成文件
This commit is contained in:
@@ -8,7 +8,15 @@
|
|||||||
build/*
|
build/*
|
||||||
build2/*
|
build2/*
|
||||||
build3/*
|
build3/*
|
||||||
|
build4/*
|
||||||
|
build_lsx/*
|
||||||
|
build_new/*
|
||||||
test/
|
test/
|
||||||
|
|
||||||
|
ZebraCrossing_Detection-master/
|
||||||
|
ZebraCrossing_Detection-master.zip
|
||||||
|
smartcar2_code.txt
|
||||||
|
smartcar_test
|
||||||
|
|
||||||
# Local History for Visual Studio Code
|
# Local History for Visual Studio Code
|
||||||
.history/
|
.history/
|
||||||
|
|||||||
Binary file not shown.
@@ -0,0 +1,123 @@
|
|||||||
|
"""
|
||||||
|
Mild Mega r2 → mild_v13.bin 导出
|
||||||
|
架构: 3→9→9→13→13→19→19→19→44→44→64→64 → head(dw,64→64) → out(64→9)
|
||||||
|
无残差连接, 纯顺序 conv_bn_relu + SE
|
||||||
|
"""
|
||||||
|
import sys, struct, os
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
class MildMegaR2(nn.Module):
|
||||||
|
def __init__(self, num_classes=4):
|
||||||
|
super().__init__()
|
||||||
|
channels = [3, 9, 9, 13, 13, 19, 19, 19, 44, 44, 64, 64]
|
||||||
|
strides = [2, 1, 2, 1, 2, 1, 1, 1, 1, 1, 1]
|
||||||
|
|
||||||
|
for i in range(11):
|
||||||
|
in_c, out_c, s = channels[i], channels[i + 1], strides[i]
|
||||||
|
if s == 2:
|
||||||
|
conv = nn.Conv2d(in_c, out_c, 3, s, 1, bias=False)
|
||||||
|
elif in_c == out_c:
|
||||||
|
conv = nn.Conv2d(in_c, out_c, 3, 1, 1, groups=in_c, bias=False)
|
||||||
|
else:
|
||||||
|
conv = nn.Conv2d(in_c, out_c, 1, 1, 0, bias=False)
|
||||||
|
|
||||||
|
setattr(self, f'b{i + 1}', nn.Sequential(
|
||||||
|
conv,
|
||||||
|
nn.BatchNorm2d(out_c),
|
||||||
|
nn.ReLU(inplace=True),
|
||||||
|
))
|
||||||
|
|
||||||
|
mid = out_c // 2
|
||||||
|
setattr(self, f'se_{i}', nn.Sequential(
|
||||||
|
nn.AdaptiveAvgPool2d(1),
|
||||||
|
nn.Conv2d(out_c, mid, 1, bias=True),
|
||||||
|
nn.ReLU(inplace=True),
|
||||||
|
nn.Conv2d(mid, out_c, 1, bias=True),
|
||||||
|
nn.Sigmoid(),
|
||||||
|
))
|
||||||
|
|
||||||
|
# head: depthwise 3x3, 64→64
|
||||||
|
self.head = nn.Sequential(
|
||||||
|
nn.Conv2d(64, 64, 3, 1, 1, groups=64, bias=False),
|
||||||
|
nn.BatchNorm2d(64),
|
||||||
|
nn.ReLU(inplace=True),
|
||||||
|
)
|
||||||
|
# output: 1x1 conv, 64→(4 class + 1 bg + 4 bbox) = 9
|
||||||
|
self.out = nn.Conv2d(64, 9, 1, bias=True)
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
for i in range(11):
|
||||||
|
x = getattr(self, f'b{i + 1}')(x)
|
||||||
|
se_out = getattr(self, f'se_{i}')(x)
|
||||||
|
x = x * se_out
|
||||||
|
x = self.head(x)
|
||||||
|
x = self.out(x)
|
||||||
|
return x
|
||||||
|
|
||||||
|
|
||||||
|
def export(pth_path, bin_path):
|
||||||
|
model = MildMegaR2(num_classes=4)
|
||||||
|
ckpt = torch.load(pth_path, map_location='cpu', weights_only=False)
|
||||||
|
|
||||||
|
# Handle wrapped checkpoint
|
||||||
|
if isinstance(ckpt, dict) and 'model' in ckpt:
|
||||||
|
ckpt = ckpt['model']
|
||||||
|
if isinstance(ckpt, dict) and 'state_dict' in ckpt:
|
||||||
|
ckpt = ckpt['state_dict']
|
||||||
|
|
||||||
|
# Remove 'module.' prefix (from DataParallel/DDP)
|
||||||
|
cleaned = {}
|
||||||
|
for k, v in ckpt.items():
|
||||||
|
nk = k.replace('module.', '')
|
||||||
|
cleaned[nk] = v
|
||||||
|
|
||||||
|
model.load_state_dict(cleaned, strict=False)
|
||||||
|
model.eval()
|
||||||
|
|
||||||
|
params = {}
|
||||||
|
for name, param in model.named_parameters():
|
||||||
|
params[name] = param.detach().cpu().numpy().astype(np.float32)
|
||||||
|
for name, buf in model.named_buffers():
|
||||||
|
params[name] = buf.detach().cpu().numpy().astype(np.float32)
|
||||||
|
|
||||||
|
total = sum(v.size for v in params.values())
|
||||||
|
print(f"Layers: {len(params)}, params: {total:,} ≈ {total / 1000:.1f}K")
|
||||||
|
|
||||||
|
for k, v in sorted(params.items()):
|
||||||
|
sh = str(list(v.shape))
|
||||||
|
print(f" {k:50s} {sh:25s} {v.size:,}")
|
||||||
|
|
||||||
|
with open(bin_path, 'wb') as f:
|
||||||
|
f.write(struct.pack('i', len(params)))
|
||||||
|
for name in sorted(params.keys()):
|
||||||
|
arr = params[name]
|
||||||
|
nb = name.encode('utf-8')
|
||||||
|
f.write(struct.pack('i', len(nb)))
|
||||||
|
f.write(nb)
|
||||||
|
f.write(struct.pack('i', arr.ndim))
|
||||||
|
for d in arr.shape:
|
||||||
|
f.write(struct.pack('i', d))
|
||||||
|
f.write(arr.tobytes())
|
||||||
|
|
||||||
|
size_kb = os.path.getsize(bin_path) / 1024
|
||||||
|
print(f"\nExported: {bin_path} ({size_kb:.0f} KB)")
|
||||||
|
|
||||||
|
# Quick inference test
|
||||||
|
x = torch.randn(1, 3, 120, 160)
|
||||||
|
with torch.no_grad():
|
||||||
|
y = model(x)
|
||||||
|
print(f"Input: {list(x.shape)} → Output: {list(y.shape)} (OK)")
|
||||||
|
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == '__main__':
|
||||||
|
import os
|
||||||
|
os.chdir(os.path.dirname(os.path.abspath(__file__)))
|
||||||
|
pth = 'best_v13_mild_r2.pth'
|
||||||
|
if not os.path.exists(pth):
|
||||||
|
print(f"ERROR: {pth} not found")
|
||||||
|
sys.exit(1)
|
||||||
|
export(pth, 'mild_v13.bin')
|
||||||
@@ -0,0 +1,18 @@
|
|||||||
|
import torch, sys, os
|
||||||
|
os.chdir(os.path.dirname(os.path.abspath(__file__)))
|
||||||
|
pth = 'best_v13_mild_r2.pth'
|
||||||
|
ckpt = torch.load(pth, map_location='cpu', weights_only=False)
|
||||||
|
print(f"Type: {type(ckpt).__name__}")
|
||||||
|
if isinstance(ckpt, dict):
|
||||||
|
if 'model' in ckpt:
|
||||||
|
print("Has 'model' key")
|
||||||
|
ckpt = ckpt['model']
|
||||||
|
if 'state_dict' in ckpt:
|
||||||
|
print("Has 'state_dict' key")
|
||||||
|
ckpt = ckpt['state_dict']
|
||||||
|
print(f"\nTotal keys: {len(ckpt)}")
|
||||||
|
for k, v in sorted(ckpt.items()):
|
||||||
|
shape = tuple(v.shape) if hasattr(v, 'shape') else 'scalar'
|
||||||
|
print(f" {k:50s} {str(shape)}")
|
||||||
|
else:
|
||||||
|
print("Not a dict, top-level keys unknown")
|
||||||
Reference in New Issue
Block a user