| """
|
| KoreanOCRNet + TrafficSignNet ์์ํ ์คํ ์คํฌ๋ฆฝํธ.
|
|
|
| Phase 1 base_W8A8.py / base_train_1bit_kd.py ํจํด ์ฌํ์ฉ.
|
| fake-quant PTQ (W8A8 / W4A16 / 1-Bit) + ONNX ๋ด๋ณด๋ด๊ธฐ + val ํ๊ฐ.
|
|
|
| ์คํ ๋งคํ:
|
| E2: FP16 ๊ฒ์ถ๊ธฐ + W8A8 ์ธ์๊ธฐ โ ๋ณธ ์คํฌ๋ฆฝํธ๋ก w8a8 ๋ด๋ณด๋ด๊ธฐ ํ ์ ํ๋ ์ธก์
|
| E3: W8A8 ์ ์ฒด โ ๊ฒ์ถ๊ธฐ(E1) + ์ธ์๊ธฐ(E2)
|
| E4-recog: W4A16 ์ธ์๊ธฐ
|
| E7-recog: 1-Bit ์ธ์๊ธฐ
|
|
|
| ์ฌ์ฉ๋ฒ:
|
| python src/quant/quantize_recognizers.py # ์ ์ฒด (w8a8/w4a16/1bit)
|
| python src/quant/quantize_recognizers.py --mode w8a8 # ํน์ ๋ชจ๋
|
| python src/quant/quantize_recognizers.py --eval_only # ๊ธฐ์กด ONNX ํ๊ฐ๋ง
|
| """
|
|
|
| import argparse
|
| import io
|
| import sys
|
| import time
|
| from pathlib import Path
|
|
|
| import numpy as np
|
| import torch
|
| import torch.nn as nn
|
|
|
|
|
| if sys.platform.startswith("win"):
|
| sys.stdout = io.TextIOWrapper(sys.stdout.buffer, encoding="utf-8", errors="replace")
|
| sys.stderr = io.TextIOWrapper(sys.stderr.buffer, encoding="utf-8", errors="replace")
|
|
|
| ROOT = Path(__file__).parent.parent.parent
|
| sys.path.insert(0, str(ROOT))
|
|
|
| MODEL_SPACE = ROOT / "model_space"
|
| MODEL_DIR = ROOT / "models"
|
| OCR_CKPT = MODEL_DIR / "korean_ocr_best.pth"
|
| TSIGN_CKPT = MODEL_SPACE / "traffic_sign_net_best.pth"
|
| OCR_DATA = ROOT / "data" / "korean_ocr"
|
| GTSDB_DIR = ROOT / "data" / "GTSDB" / "FullIJCNN2013"
|
|
|
| MODEL_SPACE.mkdir(parents=True, exist_ok=True)
|
|
|
|
|
|
|
|
|
|
|
|
|
| def _quantizable(name: str, module: nn.Module) -> bool:
|
| """Conv2d / Linear ๋ ์ด์ด๋ง ๋์ (BN, Activation ์ ์ธ)."""
|
| return isinstance(module, (nn.Conv2d, nn.Linear))
|
|
|
|
|
| def apply_w8a8_ptq(model: nn.Module) -> int:
|
| """Per-output-channel MinMax W8A8 fake-quantization (base_W8A8.py ๋์ผ ๋ฐฉ์)."""
|
| count = 0
|
| for name, module in model.named_modules():
|
| if not _quantizable(name, module):
|
| continue
|
| with torch.no_grad():
|
| w = module.weight.data
|
| if w.dim() == 4:
|
| max_val = w.view(w.size(0), -1).abs().max(dim=1)[0].view(-1, 1, 1, 1)
|
| else:
|
| max_val = w.abs().max(dim=1)[0].view(-1, 1)
|
| scale = (max_val / 127.0).clamp(min=1e-8)
|
| q_w = torch.round(w / scale).clamp(-128, 127)
|
| module.weight.data = q_w * scale
|
| count += 1
|
| return count
|
|
|
|
|
| def apply_w4a16_ptq(model: nn.Module) -> int:
|
| """Per-output-channel MinMax W4A16 fake-quantization."""
|
| count = 0
|
| for name, module in model.named_modules():
|
| if not _quantizable(name, module):
|
| continue
|
| with torch.no_grad():
|
| w = module.weight.data
|
| if w.dim() == 4:
|
| max_val = w.view(w.size(0), -1).abs().max(dim=1)[0].view(-1, 1, 1, 1)
|
| else:
|
| max_val = w.abs().max(dim=1)[0].view(-1, 1)
|
| scale = (max_val / 7.0).clamp(min=1e-8)
|
| q_w = torch.round(w / scale).clamp(-8, 7)
|
| module.weight.data = q_w * scale
|
| count += 1
|
| return count
|
|
|
|
|
| def apply_1bit_ptq(model: nn.Module) -> int:
|
| """PTB: Post-Training Binarization (sign(W) ร ||W||_1 / n, base_train_1bit_kd.py ํจํด)."""
|
| count = 0
|
| for name, module in model.named_modules():
|
| if not _quantizable(name, module):
|
| continue
|
| with torch.no_grad():
|
| w = module.weight.data
|
| if w.dim() == 4:
|
| scale = w.abs().mean(dim=(1, 2, 3), keepdim=True)
|
| else:
|
| scale = w.abs().mean(dim=1, keepdim=True)
|
| binary_w = torch.sign(w)
|
| binary_w[binary_w == 0] = 1.0
|
| module.weight.data = binary_w * scale
|
| count += 1
|
| return count
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| def export_to_onnx(
|
| model: nn.Module,
|
| dummy: torch.Tensor,
|
| out_path: Path,
|
| input_names: list,
|
| output_names: list,
|
| opset: int = 14,
|
| ):
|
| """PyTorch ๋ชจ๋ธ โ ONNX (dynamo=False)."""
|
| try:
|
| import onnxslim
|
|
|
| use_slim = True
|
| except ImportError:
|
| use_slim = False
|
|
|
| import tempfile
|
|
|
| with tempfile.NamedTemporaryFile(suffix=".onnx", delete=False) as tmp:
|
| tmp_path = Path(tmp.name)
|
|
|
| torch.onnx.export(
|
| model.cpu().eval(),
|
| dummy.cpu(),
|
| str(tmp_path),
|
| opset_version=opset,
|
| input_names=input_names,
|
| output_names=output_names,
|
| dynamic_axes={input_names[0]: {0: "batch"}, output_names[0]: {0: "batch"}},
|
| do_constant_folding=True,
|
| dynamo=False,
|
| )
|
|
|
| if use_slim:
|
| onnxslim.slim(str(tmp_path), str(out_path))
|
| tmp_path.unlink(missing_ok=True)
|
| else:
|
| tmp_path.rename(out_path)
|
|
|
| size_mb = out_path.stat().st_size / 1024 / 1024
|
| print(f" -> {out_path.name} ({size_mb:.3f} MB)")
|
| return out_path
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| def load_ocr_model() -> nn.Module:
|
| from src.korean_ocr_model import KoreanOCRNet
|
|
|
| model = KoreanOCRNet(num_classes=2350)
|
| if OCR_CKPT.exists():
|
| state = torch.load(str(OCR_CKPT), map_location="cpu", weights_only=True)
|
| model.load_state_dict(state)
|
| print(f" OCR ์ฒดํฌํฌ์ธํธ ๋ก๋: {OCR_CKPT}")
|
| else:
|
| print(f" [WARN] OCR ์ฒดํฌํฌ์ธํธ ์์: {OCR_CKPT}")
|
| return model.eval()
|
|
|
|
|
| def quantize_ocr(mode: str) -> Path:
|
| """KoreanOCRNet ์์ํ + ONNX ๋ด๋ณด๋ด๊ธฐ."""
|
| model = load_ocr_model()
|
| dummy = torch.zeros(1, 1, 64, 64)
|
|
|
| if mode == "fp32":
|
| out = MODEL_SPACE / "korean_ocr_net_fp32.onnx"
|
| elif mode == "w8a8":
|
| n = apply_w8a8_ptq(model)
|
| print(f" W8A8 fake-quant: {n} ๋ ์ด์ด")
|
| out = MODEL_SPACE / "korean_ocr_net_w8a8.onnx"
|
| elif mode == "w4a16":
|
| n = apply_w4a16_ptq(model)
|
| print(f" W4A16 fake-quant: {n} ๋ ์ด์ด")
|
| out = MODEL_SPACE / "korean_ocr_net_w4a16.onnx"
|
| elif mode == "1bit":
|
| n = apply_1bit_ptq(model)
|
| print(f" 1-Bit PTB: {n} ๋ ์ด์ด")
|
| out = MODEL_SPACE / "korean_ocr_net_1bit.onnx"
|
| else:
|
| raise ValueError(f"Unknown mode: {mode}")
|
|
|
| return export_to_onnx(model, dummy, out, ["image"], ["logits"])
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| def load_tsign_model(num_classes: int = 43) -> nn.Module:
|
| from src.model import TrafficSignNet
|
|
|
| model = TrafficSignNet(num_classes=num_classes)
|
| if TSIGN_CKPT.exists():
|
| ckpt = torch.load(str(TSIGN_CKPT), map_location="cpu", weights_only=False)
|
| nc = ckpt.get("num_classes", num_classes)
|
| if nc != num_classes:
|
| model = TrafficSignNet(num_classes=nc)
|
| model.load_state_dict(ckpt["model_state"])
|
| print(f" TrafficSignNet ์ฒดํฌํฌ์ธํธ ๋ก๋: {TSIGN_CKPT} (num_classes={nc})")
|
| else:
|
| print(f" [WARN] TrafficSignNet ์ฒดํฌํฌ์ธํธ ์์: {TSIGN_CKPT}")
|
| return model.eval()
|
|
|
|
|
| def quantize_tsign(mode: str) -> Path:
|
| """TrafficSignNet ์์ํ + ONNX ๋ด๋ณด๋ด๊ธฐ."""
|
| model = load_tsign_model()
|
| dummy = torch.zeros(1, 3, 32, 32)
|
|
|
| if mode == "fp32":
|
| out = MODEL_SPACE / "traffic_sign_net_fp32.onnx"
|
| elif mode == "w8a8":
|
| n = apply_w8a8_ptq(model)
|
| print(f" W8A8 fake-quant: {n} ๋ ์ด์ด")
|
| out = MODEL_SPACE / "traffic_sign_net_w8a8.onnx"
|
| elif mode == "w4a16":
|
| n = apply_w4a16_ptq(model)
|
| print(f" W4A16 fake-quant: {n} ๋ ์ด์ด")
|
| out = MODEL_SPACE / "traffic_sign_net_w4a16.onnx"
|
| elif mode == "1bit":
|
| n = apply_1bit_ptq(model)
|
| print(f" 1-Bit PTB: {n} ๋ ์ด์ด")
|
| out = MODEL_SPACE / "traffic_sign_net_1bit.onnx"
|
| else:
|
| raise ValueError(f"Unknown mode: {mode}")
|
|
|
| if mode != "fp32" or not out.exists():
|
| export_to_onnx(model, dummy, out, ["images"], ["logits"])
|
| return out
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| def eval_ocr_onnx(onnx_path: Path, max_samples: int = 5000) -> dict:
|
| """KoreanOCRNet ONNX ํ๊ฐ (data/korean_ocr/val/)."""
|
| import onnxruntime as ort
|
| from torchvision import datasets, transforms
|
|
|
| val_dir = OCR_DATA / "val"
|
| if not val_dir.exists():
|
| print(f" [SKIP] OCR val ๋ฐ์ดํฐ ์์: {val_dir}")
|
| return {}
|
|
|
|
|
| class NumericalImageFolder(datasets.ImageFolder):
|
| def find_classes(self, directory):
|
| import os
|
|
|
| classes = sorted([d.name for d in os.scandir(directory) if d.is_dir()], key=int)
|
| class_to_idx = {cls: int(cls) for cls in classes}
|
| return classes, class_to_idx
|
|
|
| transform = transforms.Compose(
|
| [
|
| transforms.Grayscale(num_output_channels=1),
|
| transforms.Resize((64, 64)),
|
| transforms.ToTensor(),
|
| transforms.Normalize(mean=[0.5], std=[0.5]),
|
| ]
|
| )
|
|
|
| val_ds = NumericalImageFolder(root=str(val_dir), transform=transform)
|
|
|
| indices = list(range(min(max_samples, len(val_ds))))
|
| from torch.utils.data import DataLoader, Subset
|
|
|
| subset = Subset(val_ds, indices)
|
| loader = DataLoader(subset, batch_size=256, shuffle=False, num_workers=0)
|
|
|
| sess = ort.InferenceSession(str(onnx_path), providers=["CPUExecutionProvider"])
|
| input_name = sess.get_inputs()[0].name
|
| output_name = sess.get_outputs()[0].name
|
|
|
| correct1 = correct5 = total = 0
|
| t0 = time.time()
|
|
|
| for imgs, labels in loader:
|
| imgs_np = imgs.numpy()
|
| out = sess.run([output_name], {input_name: imgs_np})[0]
|
| top5 = np.argsort(out, axis=1)[:, -5:][:, ::-1]
|
| labels_np = labels.numpy()
|
| correct1 += (top5[:, 0] == labels_np).sum()
|
| correct5 += sum(labels_np[i] in top5[i] for i in range(len(labels_np)))
|
| total += len(labels_np)
|
|
|
| elapsed = time.time() - t0
|
| top1 = correct1 / total * 100
|
| top5 = correct5 / total * 100
|
| fps = total / elapsed
|
|
|
| return {"top1": round(top1, 2), "top5": round(top5, 2), "samples": total, "fps": round(fps, 1)}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| def eval_tsign_onnx(onnx_path: Path) -> dict:
|
| """TrafficSignNet ONNX ํ๊ฐ (GTSDB val ํฌ๋กญ)."""
|
| import onnxruntime as ort
|
| from torch.utils.data import random_split
|
|
|
| from src.detect.train_traffic_sign_net import GTSDBCropDataset
|
|
|
| full_ds = GTSDBCropDataset(GTSDB_DIR, img_size=32, augment=False)
|
| n_val = max(1, int(len(full_ds) * 0.2))
|
| n_train = len(full_ds) - n_val
|
| _, val_ds = random_split(full_ds, [n_train, n_val], generator=torch.Generator().manual_seed(42))
|
|
|
|
|
| sess = ort.InferenceSession(str(onnx_path), providers=["CPUExecutionProvider"])
|
| input_name = sess.get_inputs()[0].name
|
| output_name = sess.get_outputs()[0].name
|
|
|
| _MEAN = np.array([0.485, 0.456, 0.406], dtype=np.float32)
|
| _STD = np.array([0.229, 0.224, 0.225], dtype=np.float32)
|
|
|
| correct1 = correct5 = 0
|
| t0 = time.time()
|
|
|
| for img_tensor, label in val_ds:
|
|
|
| img_np = img_tensor.numpy()[np.newaxis]
|
| out = sess.run([output_name], {input_name: img_np})[0]
|
| top5 = np.argsort(out[0])[::-1][:5]
|
| if top5[0] == label:
|
| correct1 += 1
|
| if label in top5:
|
| correct5 += 1
|
|
|
| total = len(val_ds)
|
| elapsed = time.time() - t0
|
| top1 = correct1 / total * 100
|
| top5 = correct5 / total * 100
|
|
|
| return {
|
| "top1": round(top1, 2),
|
| "top5": round(top5, 2),
|
| "samples": total,
|
| "fps": round(total / elapsed, 1),
|
| }
|
|
|
|
|
|
|
|
|
|
|
|
|
| MODES = ["fp32", "w8a8", "w4a16", "1bit"]
|
|
|
| EXP_MAP = {
|
| "fp32": "E0 ๊ธฐ์ค์ ",
|
| "w8a8": "E2/E3 W8A8",
|
| "w4a16": "E4 W4A16",
|
| "1bit": "E7 1-Bit PTB",
|
| }
|
|
|
|
|
| def run_all(
|
| modes: list[str], eval_ocr: bool = True, eval_tsign: bool = True, ocr_samples: int = 5000
|
| ):
|
| results = []
|
|
|
| for mode in modes:
|
| print(f"\n{'=' * 55}")
|
| print(f"[{mode.upper()}] KoreanOCRNet + TrafficSignNet")
|
| print(f"{'=' * 55}")
|
|
|
|
|
| print(" [OCR] ์์ํ ๋ด๋ณด๋ด๊ธฐ...")
|
| try:
|
| ocr_path = quantize_ocr(mode)
|
| except Exception as e:
|
| print(f" [OCR] ๋ด๋ณด๋ด๊ธฐ ์คํจ: {e}")
|
| ocr_path = None
|
|
|
| ocr_metrics = {}
|
| if eval_ocr and ocr_path and ocr_path.exists():
|
| print(f" [OCR] ํ๊ฐ ์ค (์ต๋ {ocr_samples}๊ฐ)...")
|
| ocr_metrics = eval_ocr_onnx(ocr_path, max_samples=ocr_samples)
|
| if ocr_metrics:
|
| print(
|
| f" Top-1={ocr_metrics['top1']:.2f}% "
|
| f"Top-5={ocr_metrics['top5']:.2f}% "
|
| f"({ocr_metrics['samples']}๊ฐ, {ocr_metrics['fps']:.0f} FPS)"
|
| )
|
|
|
|
|
| print(" [TrafficSign] ์์ํ ๋ด๋ณด๋ด๊ธฐ...")
|
| try:
|
| ts_path = quantize_tsign(mode)
|
| except Exception as e:
|
| print(f" [TrafficSign] ๋ด๋ณด๋ด๊ธฐ ์คํจ: {e}")
|
| ts_path = None
|
|
|
| ts_metrics = {}
|
| if eval_tsign and ts_path and ts_path.exists():
|
| print(" [TrafficSign] ํ๊ฐ ์ค...")
|
| ts_metrics = eval_tsign_onnx(ts_path)
|
| if ts_metrics:
|
| print(
|
| f" Top-1={ts_metrics['top1']:.2f}% "
|
| f"Top-5={ts_metrics['top5']:.2f}% "
|
| f"({ts_metrics['samples']}๊ฐ, {ts_metrics['fps']:.0f} FPS)"
|
| )
|
|
|
|
|
| ocr_size = ocr_path.stat().st_size / 1024 / 1024 if ocr_path and ocr_path.exists() else 0
|
| ts_size = ts_path.stat().st_size / 1024 / 1024 if ts_path and ts_path.exists() else 0
|
|
|
| results.append(
|
| {
|
| "mode": mode,
|
| "exp": EXP_MAP.get(mode, mode),
|
| "ocr_top1": ocr_metrics.get("top1", "โ"),
|
| "ocr_top5": ocr_metrics.get("top5", "โ"),
|
| "ts_top1": ts_metrics.get("top1", "โ"),
|
| "ts_top5": ts_metrics.get("top5", "โ"),
|
| "ocr_size": round(ocr_size, 3),
|
| "ts_size": round(ts_size, 3),
|
| }
|
| )
|
|
|
|
|
| print(f"\n{'=' * 70}")
|
| print("์ธ์๊ธฐ ์์ํ ๊ฒฐ๊ณผ ์์ฝ")
|
| print(f"{'=' * 70}")
|
| hdr = (
|
| f"{'๋ชจ๋':<8} {'์คํ':<12} "
|
| f"{'OCR Top1':>9} {'OCR Top5':>9} {'TS Top1':>8} {'TS Top5':>8} "
|
| f"{'OCR MB':>7} {'TS MB':>6}"
|
| )
|
| print(hdr)
|
| print("-" * 70)
|
|
|
| e0 = next((r for r in results if r["mode"] == "fp32"), None)
|
|
|
| for r in results:
|
| ocr1 = r["ocr_top1"]
|
| ts1 = r["ts_top1"]
|
|
|
| if e0 and r["mode"] != "fp32":
|
| if isinstance(ocr1, float) and isinstance(e0["ocr_top1"], float):
|
| ocr1_str = f"{ocr1:.1f} ({ocr1 - e0['ocr_top1']:+.1f})"
|
| else:
|
| ocr1_str = str(ocr1)
|
| if isinstance(ts1, float) and isinstance(e0["ts_top1"], float):
|
| ts1_str = f"{ts1:.1f} ({ts1 - e0['ts_top1']:+.1f})"
|
| else:
|
| ts1_str = str(ts1)
|
| else:
|
| ocr1_str = f"{ocr1:.1f}" if isinstance(ocr1, float) else str(ocr1)
|
| ts1_str = f"{ts1:.1f}" if isinstance(ts1, float) else str(ts1)
|
|
|
| print(
|
| f"{r['mode']:<8} {r['exp']:<12} "
|
| f"{ocr1_str:>9} {str(r['ocr_top5']):>9} "
|
| f"{ts1_str:>8} {str(r['ts_top5']):>8} "
|
| f"{r['ocr_size']:>7.3f} {r['ts_size']:>6.3f}"
|
| )
|
|
|
| print("\n[๋ฏผ๊ฐ๋ ๋ถ์]")
|
| if e0:
|
| for r in results:
|
| if r["mode"] == "fp32":
|
| continue
|
| for name, base_key, key in [
|
| ("OCR Top-1", "ocr_top1", "ocr_top1"),
|
| ("TrafficSign Top-1", "ts_top1", "ts_top1"),
|
| ]:
|
| bv = e0[base_key]
|
| v = r[key]
|
| if isinstance(v, float) and isinstance(bv, float):
|
| d = v - bv
|
| p = d / bv * 100
|
| print(f" {r['mode']} {name}: {bv:.1f}% -> {v:.1f}% ({d:+.1f}pp, {p:+.1f}%)")
|
|
|
| print("\n์ ๊ฒฐ๊ณผ๋ฅผ docs/EXPERIMENTS.md ์ธ์ ๊ฒฐ๊ณผ ํ์ ๊ธฐ์
ํ์ธ์.")
|
| return results
|
|
|
|
|
| def main():
|
| parser = argparse.ArgumentParser(description="์ธ์๊ธฐ ์์ํ + ํ๊ฐ")
|
| parser.add_argument("--mode", choices=MODES + ["all"], default="all")
|
| parser.add_argument("--no_ocr", action="store_true", help="OCR ํ๊ฐ ๊ฑด๋๋ฐ๊ธฐ")
|
| parser.add_argument("--no_tsign", action="store_true", help="TrafficSign ํ๊ฐ ๊ฑด๋๋ฐ๊ธฐ")
|
| parser.add_argument(
|
| "--ocr_samples", type=int, default=5000, help="OCR val ์ต๋ ์ํ ์ (๊ธฐ๋ณธ 5000)"
|
| )
|
| args = parser.parse_args()
|
|
|
| modes = MODES if args.mode == "all" else [args.mode]
|
| run_all(
|
| modes, eval_ocr=not args.no_ocr, eval_tsign=not args.no_tsign, ocr_samples=args.ocr_samples
|
| )
|
|
|
|
|
| if __name__ == "__main__":
|
| main()
|
|
|