# -*- coding: utf-8 -*- """prepare_train_data_v3.py - 新数据集88407张:上半区全量训练数据 + 下半区全量裁剪(供挖矿)。 输入: data/gallery_images + data/train_meta_88407.json (fname=url basename) 输出: upper_data/{images, train_groups.json, val_groups.json, manifest.csv} # 上半区全量, name+no分组 lower_all_data/images/ # 下半区全量裁剪, 挖矿原料 分组: card_name_ch + card_no, multi组10%按组held-out, 单张组全进train。 """ import argparse import csv import json import os import random import sys import time from collections import defaultdict import numpy as np from PIL import Image IMG_SIZE = 392 HALF = 196 def letterbox(pil_img): img = pil_img.convert("RGB") w, h = img.size r = min(IMG_SIZE / h, IMG_SIZE / w) nw, nh = int(w * r), int(h * r) img = img.resize((nw, nh), Image.BICUBIC) canvas = Image.new("RGB", (IMG_SIZE, IMG_SIZE), (0, 0, 0)) canvas.paste(img, ((IMG_SIZE - nw) // 2, (IMG_SIZE - nh) // 2)) return canvas def main(): ap = argparse.ArgumentParser() ap.add_argument("--wzj", default="/home/user/顾工交接/wzj") ap.add_argument("--meta", default="/home/user/顾工交接/wzj/data/train_meta_88407.json") ap.add_argument("--gallery", default="/home/user/顾工交接/wzj/data/gallery_images") ap.add_argument("--upper-out", default="/home/user/顾工交接/wzj/upper_data") ap.add_argument("--lower-out", default="/home/user/顾工交接/wzj/lower_all_data") ap.add_argument("--val-ratio", type=float, default=0.1) ap.add_argument("--seed", type=int, default=42) ap.add_argument("--limit", type=int, default=0) ap.add_argument("--shard", type=int, default=0, help="0=单进程; >=1 表示第shard片(1-based)") ap.add_argument("--nshards", type=int, default=1) args = ap.parse_args() sys.path.insert(0, args.wzj) os.chdir(args.wzj) import config from modules.yolo_detector import CardDetector with open(args.meta, "r", encoding="utf-8") as f: meta = json.load(f) records = meta["records"] if args.limit: records = records[: args.limit] print(f"[meta] records={len(records)}", flush=True) items = [] miss = 0 for r in records: path = os.path.join(args.gallery, r["fname"]) if not (os.path.exists(path) and os.path.getsize(path) > 0): miss += 1 continue items.append({**r, "path": path}) print(f"[paths] ok={len(items)} missing={miss}", flush=True) os.makedirs(os.path.join(args.upper_out, "images"), exist_ok=True) os.makedirs(os.path.join(args.lower_out, "images"), exist_ok=True) print(f"[yolo] load {config.YOLO_MODEL_PATH}", flush=True) det = CardDetector(config.YOLO_MODEL_PATH) # 分片支持(多进程并行裁剪时用) shard_items = items if args.nshards > 1: shard_items = [it for i, it in enumerate(items) if i % args.nshards == (args.shard - 1)] print(f"[shard] {args.shard}/{args.nshards} items={len(shard_items)}", flush=True) n_ok = n_skip = n_fail = 0 t0 = time.time() prepared = [] fail_list = [] for i, job in enumerate(shard_items): sid = job["sample_id"] up_p = os.path.join(args.upper_out, "images", f"{sid}.jpg") lo_p = os.path.join(args.lower_out, "images", f"{sid}.jpg") if (os.path.exists(up_p) and os.path.getsize(up_p) > 0 and os.path.exists(lo_p) and os.path.getsize(lo_p) > 0): n_skip += 1 prepared.append(job) continue try: img = Image.open(job["path"]).convert("RGB") arr = np.array(img) crop = det.detect_and_crop(arr) if crop is None: crop = arr lb = letterbox(Image.fromarray(crop)) lb.crop((0, 0, IMG_SIZE, HALF)).save(up_p, format="JPEG", quality=95) lb.crop((0, HALF, IMG_SIZE, IMG_SIZE)).save(lo_p, format="JPEG", quality=95) n_ok += 1 prepared.append(job) except Exception as e: n_fail += 1 fail_list.append((sid, str(e)[:100])) if n_fail <= 20: print(f" [FAIL] {sid}: {e}", flush=True) if (i + 1) % 500 == 0 or (i + 1) == len(shard_items): print(f" crop {i+1}/{len(shard_items)} ok={n_ok} skip={n_skip} fail={n_fail} " f"elapsed={(time.time()-t0)/60:.1f}m", flush=True) print(f"[crop done] ok={n_ok} skip={n_skip} fail={n_fail} prepared={len(prepared)}", flush=True) # 分组只由主分片(单进程)产出,避免多进程写冲突 if args.nshards == 1: groups = defaultdict(list) for r in prepared: ok = True for p in (os.path.join(args.upper_out, "images", f"{r['sample_id']}.jpg"),): if not (os.path.exists(p) and os.path.getsize(p) > 0): ok = False if not ok: continue groups[(r.get("card_name_ch") or "", r.get("card_no") or "")].append(r) multi_keys = [k for k, v in groups.items() if len(v) >= 2] single_keys = [k for k, v in groups.items() if len(v) == 1] rng = random.Random(args.seed) rng.shuffle(multi_keys) n_val = int(len(multi_keys) * args.val_ratio) val_keys = set(multi_keys[:n_val]) train_keys = set(multi_keys[n_val:]) | set(single_keys) def dump_groups(key_set, path): out = [] gid = 0 for k in sorted(key_set): out.append({ "group_id": gid, "name": k[0], "card_no": k[1], "card_ids": [m["sample_id"] for m in groups[k]], }) gid += 1 with open(path, "w", encoding="utf-8") as f: json.dump(out, f, ensure_ascii=False) return len(out), sum(len(g["card_ids"]) for g in out) n_tr_g, n_tr_s = dump_groups(train_keys, os.path.join(args.upper_out, "train_groups.json")) n_va_g, n_va_s = dump_groups(val_keys, os.path.join(args.upper_out, "val_groups.json")) split_of = {} for k in train_keys: for m in groups[k]: split_of[m["sample_id"]] = "train" for k in val_keys: for m in groups[k]: split_of[m["sample_id"]] = "val" with open(os.path.join(args.upper_out, "manifest.csv"), "w", encoding="utf-8-sig", newline="") as f: w = csv.writer(f) w.writerow(["sample_id", "card_id", "card_name_ch", "card_no", "language", "pg_label", "split"]) for m in prepared: sid = m["sample_id"] if sid not in split_of: continue w.writerow([sid, m.get("card_id") or "", m.get("card_name_ch") or "", m.get("card_no") or "", m.get("language") or "", m.get("pg_label") or "", split_of[sid]]) print(f"[upper] train_groups={n_tr_g} train_samples={n_tr_s} " f"val_groups={n_va_g} val_samples={n_va_s} -> {args.upper_out}", flush=True) stats = { "prepared": len(prepared), "groups_total": len(groups), "multi_groups": len(multi_keys), "crop_fail": n_fail, "fails": fail_list[:50], } with open(os.path.join(args.upper_out, "prepare_stats.json"), "w", encoding="utf-8") as f: json.dump(stats, f, ensure_ascii=False, indent=2) print("[DONE]", json.dumps({k: stats[k] for k in ("prepared", "groups_total", "multi_groups", "crop_fail")}, ensure_ascii=False), flush=True) else: print(f"[shard {args.shard}] done, groups skipped (main shard only)", flush=True) if __name__ == "__main__": main()