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/*
|
||||
build2/*
|
||||
build3/*
|
||||
build4/*
|
||||
build_lsx/*
|
||||
build_new/*
|
||||
test/
|
||||
|
||||
ZebraCrossing_Detection-master/
|
||||
ZebraCrossing_Detection-master.zip
|
||||
smartcar2_code.txt
|
||||
smartcar_test
|
||||
|
||||
# Local History for Visual Studio Code
|
||||
.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