| 14 | |
| 15 | # JSON 数据加载器 (确保精度和 Go 完全一致) |
| 16 | class ShadowJSONDataset(Dataset): |
| 17 | def __init__(self, json_path): |
| 18 | print(f"[*] 正在加载 JSON 数据集: {json_path} ...") |
| 19 | with open(json_path, 'r') as f: |
| 20 | self.data = json.load(f) |
| 21 | |
| 22 | def __len__(self): |
| 23 | return len(self.data) |
| 24 | |
| 25 | def __getitem__(self, idx): |
| 26 | item = self.data[idx] |
| 27 | img_key = 'image' if 'image' in item else 'Image' |
| 28 | lbl_key = 'label' if 'label' in item else 'Label' |
| 29 | |
| 30 | # 🚨【终极修复点】:将写死的 'Image' 和 'Label' 替换成动态探测出来的 img_key 和 lbl_key! |
| 31 | # 这样无论数据包是新版(小写)还是旧版(大写),均可 100% 完美自动适配,绝不报 KeyError! |
| 32 | image_tensor = torch.tensor(item[img_key], dtype=torch.float32).view(3, 32, 32) |
| 33 | label_tensor = torch.tensor(item[lbl_key], dtype=torch.long) |
| 34 | |
| 35 | return image_tensor, label_tensor |
| 36 | |
| 37 | # 增加准确率统计逻辑 |
| 38 | def get_losses_and_acc(model, loader): |