| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193 |
- # -*- 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()
|