prepare_train_data_v3.py 7.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193
  1. # -*- coding: utf-8 -*-
  2. """prepare_train_data_v3.py - 新数据集88407张:上半区全量训练数据 + 下半区全量裁剪(供挖矿)。
  3. 输入: data/gallery_images + data/train_meta_88407.json (fname=url basename)
  4. 输出:
  5. upper_data/{images, train_groups.json, val_groups.json, manifest.csv} # 上半区全量, name+no分组
  6. lower_all_data/images/ # 下半区全量裁剪, 挖矿原料
  7. 分组: card_name_ch + card_no, multi组10%按组held-out, 单张组全进train。
  8. """
  9. import argparse
  10. import csv
  11. import json
  12. import os
  13. import random
  14. import sys
  15. import time
  16. from collections import defaultdict
  17. import numpy as np
  18. from PIL import Image
  19. IMG_SIZE = 392
  20. HALF = 196
  21. def letterbox(pil_img):
  22. img = pil_img.convert("RGB")
  23. w, h = img.size
  24. r = min(IMG_SIZE / h, IMG_SIZE / w)
  25. nw, nh = int(w * r), int(h * r)
  26. img = img.resize((nw, nh), Image.BICUBIC)
  27. canvas = Image.new("RGB", (IMG_SIZE, IMG_SIZE), (0, 0, 0))
  28. canvas.paste(img, ((IMG_SIZE - nw) // 2, (IMG_SIZE - nh) // 2))
  29. return canvas
  30. def main():
  31. ap = argparse.ArgumentParser()
  32. ap.add_argument("--wzj", default="/home/user/顾工交接/wzj")
  33. ap.add_argument("--meta", default="/home/user/顾工交接/wzj/data/train_meta_88407.json")
  34. ap.add_argument("--gallery", default="/home/user/顾工交接/wzj/data/gallery_images")
  35. ap.add_argument("--upper-out", default="/home/user/顾工交接/wzj/upper_data")
  36. ap.add_argument("--lower-out", default="/home/user/顾工交接/wzj/lower_all_data")
  37. ap.add_argument("--val-ratio", type=float, default=0.1)
  38. ap.add_argument("--seed", type=int, default=42)
  39. ap.add_argument("--limit", type=int, default=0)
  40. ap.add_argument("--shard", type=int, default=0, help="0=单进程; >=1 表示第shard片(1-based)")
  41. ap.add_argument("--nshards", type=int, default=1)
  42. args = ap.parse_args()
  43. sys.path.insert(0, args.wzj)
  44. os.chdir(args.wzj)
  45. import config
  46. from modules.yolo_detector import CardDetector
  47. with open(args.meta, "r", encoding="utf-8") as f:
  48. meta = json.load(f)
  49. records = meta["records"]
  50. if args.limit:
  51. records = records[: args.limit]
  52. print(f"[meta] records={len(records)}", flush=True)
  53. items = []
  54. miss = 0
  55. for r in records:
  56. path = os.path.join(args.gallery, r["fname"])
  57. if not (os.path.exists(path) and os.path.getsize(path) > 0):
  58. miss += 1
  59. continue
  60. items.append({**r, "path": path})
  61. print(f"[paths] ok={len(items)} missing={miss}", flush=True)
  62. os.makedirs(os.path.join(args.upper_out, "images"), exist_ok=True)
  63. os.makedirs(os.path.join(args.lower_out, "images"), exist_ok=True)
  64. print(f"[yolo] load {config.YOLO_MODEL_PATH}", flush=True)
  65. det = CardDetector(config.YOLO_MODEL_PATH)
  66. # 分片支持(多进程并行裁剪时用)
  67. shard_items = items
  68. if args.nshards > 1:
  69. shard_items = [it for i, it in enumerate(items) if i % args.nshards == (args.shard - 1)]
  70. print(f"[shard] {args.shard}/{args.nshards} items={len(shard_items)}", flush=True)
  71. n_ok = n_skip = n_fail = 0
  72. t0 = time.time()
  73. prepared = []
  74. fail_list = []
  75. for i, job in enumerate(shard_items):
  76. sid = job["sample_id"]
  77. up_p = os.path.join(args.upper_out, "images", f"{sid}.jpg")
  78. lo_p = os.path.join(args.lower_out, "images", f"{sid}.jpg")
  79. if (os.path.exists(up_p) and os.path.getsize(up_p) > 0
  80. and os.path.exists(lo_p) and os.path.getsize(lo_p) > 0):
  81. n_skip += 1
  82. prepared.append(job)
  83. continue
  84. try:
  85. img = Image.open(job["path"]).convert("RGB")
  86. arr = np.array(img)
  87. crop = det.detect_and_crop(arr)
  88. if crop is None:
  89. crop = arr
  90. lb = letterbox(Image.fromarray(crop))
  91. lb.crop((0, 0, IMG_SIZE, HALF)).save(up_p, format="JPEG", quality=95)
  92. lb.crop((0, HALF, IMG_SIZE, IMG_SIZE)).save(lo_p, format="JPEG", quality=95)
  93. n_ok += 1
  94. prepared.append(job)
  95. except Exception as e:
  96. n_fail += 1
  97. fail_list.append((sid, str(e)[:100]))
  98. if n_fail <= 20:
  99. print(f" [FAIL] {sid}: {e}", flush=True)
  100. if (i + 1) % 500 == 0 or (i + 1) == len(shard_items):
  101. print(f" crop {i+1}/{len(shard_items)} ok={n_ok} skip={n_skip} fail={n_fail} "
  102. f"elapsed={(time.time()-t0)/60:.1f}m", flush=True)
  103. print(f"[crop done] ok={n_ok} skip={n_skip} fail={n_fail} prepared={len(prepared)}", flush=True)
  104. # 分组只由主分片(单进程)产出,避免多进程写冲突
  105. if args.nshards == 1:
  106. groups = defaultdict(list)
  107. for r in prepared:
  108. ok = True
  109. for p in (os.path.join(args.upper_out, "images", f"{r['sample_id']}.jpg"),):
  110. if not (os.path.exists(p) and os.path.getsize(p) > 0):
  111. ok = False
  112. if not ok:
  113. continue
  114. groups[(r.get("card_name_ch") or "", r.get("card_no") or "")].append(r)
  115. multi_keys = [k for k, v in groups.items() if len(v) >= 2]
  116. single_keys = [k for k, v in groups.items() if len(v) == 1]
  117. rng = random.Random(args.seed)
  118. rng.shuffle(multi_keys)
  119. n_val = int(len(multi_keys) * args.val_ratio)
  120. val_keys = set(multi_keys[:n_val])
  121. train_keys = set(multi_keys[n_val:]) | set(single_keys)
  122. def dump_groups(key_set, path):
  123. out = []
  124. gid = 0
  125. for k in sorted(key_set):
  126. out.append({
  127. "group_id": gid,
  128. "name": k[0],
  129. "card_no": k[1],
  130. "card_ids": [m["sample_id"] for m in groups[k]],
  131. })
  132. gid += 1
  133. with open(path, "w", encoding="utf-8") as f:
  134. json.dump(out, f, ensure_ascii=False)
  135. return len(out), sum(len(g["card_ids"]) for g in out)
  136. n_tr_g, n_tr_s = dump_groups(train_keys, os.path.join(args.upper_out, "train_groups.json"))
  137. n_va_g, n_va_s = dump_groups(val_keys, os.path.join(args.upper_out, "val_groups.json"))
  138. split_of = {}
  139. for k in train_keys:
  140. for m in groups[k]:
  141. split_of[m["sample_id"]] = "train"
  142. for k in val_keys:
  143. for m in groups[k]:
  144. split_of[m["sample_id"]] = "val"
  145. with open(os.path.join(args.upper_out, "manifest.csv"), "w", encoding="utf-8-sig", newline="") as f:
  146. w = csv.writer(f)
  147. w.writerow(["sample_id", "card_id", "card_name_ch", "card_no", "language", "pg_label", "split"])
  148. for m in prepared:
  149. sid = m["sample_id"]
  150. if sid not in split_of:
  151. continue
  152. w.writerow([sid, m.get("card_id") or "", m.get("card_name_ch") or "",
  153. m.get("card_no") or "", m.get("language") or "",
  154. m.get("pg_label") or "", split_of[sid]])
  155. print(f"[upper] train_groups={n_tr_g} train_samples={n_tr_s} "
  156. f"val_groups={n_va_g} val_samples={n_va_s} -> {args.upper_out}", flush=True)
  157. stats = {
  158. "prepared": len(prepared), "groups_total": len(groups),
  159. "multi_groups": len(multi_keys), "crop_fail": n_fail,
  160. "fails": fail_list[:50],
  161. }
  162. with open(os.path.join(args.upper_out, "prepare_stats.json"), "w", encoding="utf-8") as f:
  163. json.dump(stats, f, ensure_ascii=False, indent=2)
  164. print("[DONE]", json.dumps({k: stats[k] for k in ("prepared", "groups_total", "multi_groups", "crop_fail")}, ensure_ascii=False), flush=True)
  165. else:
  166. print(f"[shard {args.shard}] done, groups skipped (main shard only)", flush=True)
  167. if __name__ == "__main__":
  168. main()