diff --git a/.gitignore b/.gitignore index cacd9aa..91c8a27 100644 --- a/.gitignore +++ b/.gitignore @@ -123,3 +123,9 @@ logs/ # local process state for start/stop scripts scripts/.run/ + +# Train framework local runs (#95) +train/runs/ +train/data_cache/ +train/.venv/ + diff --git a/README.md b/README.md index 232a592..b1c1216 100644 --- a/README.md +++ b/README.md @@ -62,6 +62,8 @@ L5 音符节点 音高数字 OCR(+几何兜底)+ 时值线 / 高低音点 桌面在结构模式下可叠图查看 L1–L5,L3 以分割线编辑为主。 完整说明:[architecture-structure-first.md](./docs/architecture-structure-first.md) · [l3-split-model.md](./docs/l3-split-model.md) · [architecture.md](./docs/architecture.md)。 +L1–L3 **布局训练**(#92–#95):数据规范 [docs/train/l1-l3-data-spec.md](./docs/train/l1-l3-data-spec.md) · 模型方案 [l1-l3-model-design.md](./docs/train/l1-l3-model-design.md) · Framework [`train/`](./train/)。 + --- ## 仓库结构 diff --git a/docs/train/l1-l3-model-design.md b/docs/train/l1-l3-model-design.md index 173668d..7094f96 100644 --- a/docs/train/l1-l3-model-design.md +++ b/docs/train/l1-l3-model-design.md @@ -2,7 +2,7 @@ > 状态:设计定稿 v0.1(#94 / 父任务 #92) > 数据契约:[`l1-l3-data-spec.md`](./l1-l3-data-spec.md)(`layout_schema_version` **0.1**) -> 实现骨架:训练 Framework **#95**;core 推理插件为 **P2**(另开 Issue) +> 实现骨架:训练 Framework **#95** 见仓库 [`train/`](../../train/);core 推理插件为 **P2** --- @@ -264,21 +264,23 @@ learned_l1l3: --- -## 8. 训练 Framework 接口预期(给 #95) +## 8. 训练 Framework(#95 已落地) -最小可实现切片(与 #95 验收对齐): +目录:`train/`(说明见 [`train/README.md`](../../train/README.md))。 | 模块 | 职责 | |------|------| -| `data/` | 读 data-spec;collate;可视化图+框+竖线 | -| `models/` | L2 det ± L3 heat(L1 可选) | -| `metrics/` | L2 IoU/mAP;L3 count + mean_abs_x | -| `engine/` | train/eval loop、ckpt | -| `export/` | state_dict / ONNX + README「core 如何加载」段落 | -| `configs/mvp_l2_l3.yaml` | 任务开关、路径、超参 | - -一条命令:`python scripts/train.py --config configs/mvp_l2_l3.yaml` -数据:≥1 真实样本 + 合成/增强;至少 1 个 epoch 不报错。 +| `enpu_train/data/` | data-spec Dataset + 合成样本 + 可视化 | +| `enpu_train/models/` | L2 page **y-heat** + L3 row **x-heat**(轻量 CNN) | +| `enpu_train/metrics/` | L2 条带 IoU;L3 count + mean_abs_x | +| `enpu_train/engine/` | train/eval、ckpt | +| `enpu_train/export/` | state_dict(+ 可选 ONNX) | +| `configs/mvp_l2_l3.yaml` | 任务与超参 | + +```powershell +cd train +python scripts/train.py --config configs/mvp_l2_l3.yaml +``` --- diff --git a/train/README.md b/train/README.md new file mode 100644 index 0000000..13592fa --- /dev/null +++ b/train/README.md @@ -0,0 +1,121 @@ +# EnPu Train — L1–L3 布局训练 Framework(#95) + +独立于 `desktop/` 的训练应用:读 **layout data-spec** 样本,训练轻量 **L2 + L3** 模型(#94 MVP-1),导出权重。 + +| 文档 | 链接 | +|------|------| +| 数据规范 #93 | [docs/train/l1-l3-data-spec.md](../docs/train/l1-l3-data-spec.md) | +| 模型方案 #94 | [docs/train/l1-l3-model-design.md](../docs/train/l1-l3-model-design.md) | +| 样例数据 | [samples/layout/](../samples/layout/) | + +## 目录 + +```text +train/ + README.md + requirements.txt + configs/mvp_l2_l3.yaml + enpu_train/ + data/ # Dataset + 合成样本 + models/ # L2 page y-heat + L3 row x-heat + losses/ + metrics/ # L2 IoU + L3 split count / mean_abs_x + engine/ # train / eval + export/ # state_dict + ONNX + viz.py + scripts/ + train.py + eval.py + viz_sample.py + export_from_enpu_project.py + tests/ +``` + +## 环境 + +```powershell +cd train +python -m venv .venv +.\.venv\Scripts\Activate.ps1 +pip install -r requirements.txt +``` + +需要 **Python 3.10+**、**PyTorch**。有 GPU 时可在配置里设 `train.device: cuda`。 + +## 数据准备 + +1. **真实工程** → layout 样本(#93): + +```powershell +# 仓库根 +$env:PYTHONPATH = ".\core" +python scripts\export_layout_gt.py --project path\to\song.enpu.json --out samples\layout\L00x + +# 或在 train/ 下 +python scripts\export_from_enpu_project.py -p path\to\song.enpu.json -o ..\samples\layout\L00x +``` + +2. **仓库已有**:`samples/layout/L001_zuozai_baozuo/` +3. **合成**:`train.py` 会按配置自动生成 `data_cache/synth/S00*`(默认 8 张) + +抽查可视化: + +```powershell +cd train +python scripts\viz_sample.py ..\samples\layout\L001_zuozai_baozuo +``` + +## 一条命令 toy 训练 + +```powershell +cd train +python scripts\train.py --config configs\mvp_l2_l3.yaml +``` + +默认:CPU、3 epoch、真实 layout + 合成数据、写出: + +- `runs/mvp_l2_l3/last.pt` / `best.pt` +- `runs/mvp_l2_l3/history.json` +- `runs/mvp_l2_l3/export/layout_net.pt` +- `runs/mvp_l2_l3/export/onnx/*.onnx`(若 ONNX 导出可用) + +评估: + +```powershell +python scripts\eval.py --ckpt runs\mvp_l2_l3\best.pt --data ..\samples\layout +``` + +指标:`l2_mean_iou`、`l3_mean_abs_x_error`、`l3_split_count_mae` / `exact`(与 #86 线级语义同构)。 + +## 模型与 structure 对应 + +| 模型输出 | 恩谱 structure / IR | +|----------|---------------------| +| L2 1D y 热力峰值 → 水平条带框 | `items[]` layer=L2 `kind=system`;`StaffSystem.rect` | +| L3 行 crop 上 x 热力峰值 | `barlines[]` / `StaffSystem.splits`(interior x) | +| 后处理 `normalize_splits` + `splits_to_measures` | L3 `measure_derived` 框(派生,非主学) | + +**core 如何加载(P2 草图,本仓库尚未接插件):** + +1. 读 `export/layout_net.pt` 或 ONNX +2. 全图 → L2 峰值 → systems +3. 每行 crop → L3 峰值 → splits(全图 x) +4. 填入 `StructureDebug` / `PageLayout`,再跑现有 L4–L5 或仅叠图 +5. 配置预留:`ENPU_STRUCTURE_ENGINE=learned_l1l3`(实现见后续 Issue) + +规则几何管线保持 fallback。 + +## 测试 + +```powershell +cd train +python -m pytest tests -q +``` + +## 非目标 + +- 大规模真实集 / 完整合成流水线 +- core 内完整 `learned_l1l3` 推理插件 +- 精度超过几何基线的承诺 + +父任务:[Issue #92](https://github.com/loootte/EnPu/issues/92) · 本任务 [#95](https://github.com/loootte/EnPu/issues/95) diff --git a/train/configs/mvp_l2_l3.yaml b/train/configs/mvp_l2_l3.yaml new file mode 100644 index 0000000..b40b872 --- /dev/null +++ b/train/configs/mvp_l2_l3.yaml @@ -0,0 +1,32 @@ +# MVP-1: L2 systems (page y-heat) + L3 splits (row x-heat) — #95 / #94 +tasks: [l2, l3] + +data: + # real data-spec samples (#93) + roots: + - ../samples/layout + # synthetic pages generated under train/data_cache/synth + synth_count: 8 + synth_dir: data_cache/synth + page_size: [384, 512] # H, W + row_size: [64, 256] + l2_heat_len: 128 + l3_heat_len: 128 + augment: true + val_ratio: 0.25 + +train: + epochs: 3 + batch_size: 2 + lr: 0.001 + weight_decay: 0.0001 + l2_loss_weight: 1.0 + l3_loss_weight: 1.5 + device: cpu # set cuda if available + num_workers: 0 + out_dir: runs/mvp_l2_l3 + log_every: 1 + +export: + state_dict: runs/mvp_l2_l3/export/layout_net.pt + onnx_dir: runs/mvp_l2_l3/export/onnx diff --git a/train/enpu_train/__init__.py b/train/enpu_train/__init__.py new file mode 100644 index 0000000..1985f40 --- /dev/null +++ b/train/enpu_train/__init__.py @@ -0,0 +1,3 @@ +"""EnPu L1–L3 layout training framework (#95 / #92).""" + +__version__ = "0.1.0" diff --git a/train/enpu_train/data/__init__.py b/train/enpu_train/data/__init__.py new file mode 100644 index 0000000..c425bf3 --- /dev/null +++ b/train/enpu_train/data/__init__.py @@ -0,0 +1,4 @@ +from enpu_train.data.dataset import LayoutDataset, collate_layout +from enpu_train.data.synthetic import make_synthetic_layout_sample + +__all__ = ["LayoutDataset", "collate_layout", "make_synthetic_layout_sample"] diff --git a/train/enpu_train/data/dataset.py b/train/enpu_train/data/dataset.py new file mode 100644 index 0000000..4081f59 --- /dev/null +++ b/train/enpu_train/data/dataset.py @@ -0,0 +1,344 @@ +"""Dataset for layout_schema_version 0.1 samples (#95 / #93).""" + +from __future__ import annotations + +import json +import random +from pathlib import Path +from typing import Any + +import numpy as np +import torch +from PIL import Image +from torch.utils.data import Dataset + + +def _load_layout(path: Path) -> dict[str, Any]: + return json.loads(path.read_text(encoding="utf-8")) + + +def discover_samples(root: str | Path) -> list[Path]: + """Find directories containing layout.json under root.""" + root = Path(root) + if not root.is_dir(): + return [] + found = sorted({p.parent for p in root.rglob("layout.json")}) + return found + + +def build_l2_y_heat( + systems: list[dict[str, Any]], + *, + height: int, + heat_len: int, + sigma: float = 2.0, +) -> np.ndarray: + """1D heatmap along image height: peaks at system vertical centers.""" + heat = np.zeros(heat_len, dtype=np.float32) + if height <= 0: + return heat + for s in systems: + b = s.get("bbox") or s.get("box") or {} + cy = 0.5 * (float(b["y1"]) + float(b["y2"])) + idx = cy / height * (heat_len - 1) + for t in range(heat_len): + heat[t] = max(heat[t], float(np.exp(-0.5 * ((t - idx) / sigma) ** 2))) + return heat + + +def build_l3_x_heat( + splits: list[dict[str, Any] | float], + *, + x_left: float, + x_right: float, + heat_len: int, + sigma: float = 1.5, +) -> np.ndarray: + """1D heatmap along row width for interior splits (relative to crop).""" + heat = np.zeros(heat_len, dtype=np.float32) + width = max(1e-3, x_right - x_left) + for sp in splits: + x = float(sp["x"] if isinstance(sp, dict) else sp) + # relative position in [0, 1] + rel = (x - x_left) / width + if rel <= 0.0 or rel >= 1.0: + continue + idx = rel * (heat_len - 1) + for t in range(heat_len): + heat[t] = max(heat[t], float(np.exp(-0.5 * ((t - idx) / sigma) ** 2))) + return np.clip(heat, 0.0, 1.0) + + +def decode_peaks( + heat: np.ndarray, + *, + min_prominence: float = 0.3, + min_gap: int = 4, +) -> list[int]: + """Simple 1D NMS peak indices.""" + h = heat.astype(np.float32) + peaks: list[tuple[float, int]] = [] + for i in range(1, len(h) - 1): + if h[i] >= min_prominence and h[i] >= h[i - 1] and h[i] >= h[i + 1]: + peaks.append((float(h[i]), i)) + peaks.sort(reverse=True) + chosen: list[int] = [] + for _, i in peaks: + if all(abs(i - j) >= min_gap for j in chosen): + chosen.append(i) + return sorted(chosen) + + +class LayoutDataset(Dataset): + """Load enpu-layout-gt samples; produce page + per-row L3 crops.""" + + def __init__( + self, + roots: list[str | Path] | str | Path, + *, + page_size: tuple[int, int] = (384, 512), # (H, W) + row_size: tuple[int, int] = (64, 256), # (H, W) L3 crop + l2_heat_len: int = 128, + l3_heat_len: int = 128, + tasks: tuple[str, ...] = ("l2", "l3"), + augment: bool = False, + max_rows_per_page: int = 12, + seed: int = 0, + ) -> None: + if isinstance(roots, (str, Path)): + roots = [roots] + self.samples: list[Path] = [] + for r in roots: + self.samples.extend(discover_samples(r)) + if not self.samples: + raise FileNotFoundError(f"no layout.json under {roots}") + + self.page_h, self.page_w = page_size + self.row_h, self.row_w = row_size + self.l2_heat_len = l2_heat_len + self.l3_heat_len = l3_heat_len + self.tasks = tuple(tasks) + self.augment = augment + self.max_rows = max_rows_per_page + self.rng = random.Random(seed) + + def __len__(self) -> int: + return len(self.samples) + + def _load_image(self, sample_dir: Path, layout: dict[str, Any]) -> Image.Image: + rel = (layout.get("image") or {}).get("path") or "image.png" + path = sample_dir / rel + if not path.is_file(): + # fallback any image + for cand in sample_dir.glob("image.*"): + path = cand + break + img = Image.open(path).convert("RGB") + return img + + def __getitem__(self, index: int) -> dict[str, Any]: + sample_dir = self.samples[index] + layout = _load_layout(sample_dir / "layout.json") + img = self._load_image(sample_dir, layout) + orig_w, orig_h = img.size + + # optional light augment (no horizontal flip) + if self.augment: + if self.rng.random() < 0.5: + # brightness + arr = np.asarray(img).astype(np.float32) + factor = self.rng.uniform(0.85, 1.15) + arr = np.clip(arr * factor, 0, 255).astype(np.uint8) + img = Image.fromarray(arr) + if self.rng.random() < 0.3: + # slight scale via resize jitter later + pass + + # resize page + page = img.resize((self.page_w, self.page_h), Image.BILINEAR) + page_t = ( + torch.from_numpy(np.array(page, copy=True).transpose(2, 0, 1)).float() + / 255.0 + ) + + systems = list((layout.get("l2") or {}).get("systems") or []) + sx = self.page_w / max(1, orig_w) + sy = self.page_h / max(1, orig_h) + + # L2 boxes in page pixel space + l2_boxes = [] + for s in systems: + b = s.get("bbox") or s.get("box") or {} + l2_boxes.append( + [ + float(b["x1"]) * sx, + float(b["y1"]) * sy, + float(b["x2"]) * sx, + float(b["y2"]) * sy, + ] + ) + # L2 1D heat along page height (page-space box centers) + l2_heat = np.zeros(self.l2_heat_len, dtype=np.float32) + for box in l2_boxes: + cy = 0.5 * (box[1] + box[3]) + idx = cy / self.page_h * (self.l2_heat_len - 1) + sigma = 2.0 + for t in range(self.l2_heat_len): + l2_heat[t] = max( + l2_heat[t], float(np.exp(-0.5 * ((t - idx) / sigma) ** 2)) + ) + + # L3 rows + rows_meta = list((layout.get("l3") or {}).get("rows") or []) + # map system_id -> bbox from systems + sys_by_id = {str(s.get("id")): s for s in systems} + sys_by_index = { + int(s.get("index", i)): s for i, s in enumerate(systems) + } + + row_images: list[torch.Tensor] = [] + row_heats: list[torch.Tensor] = [] + row_meta: list[dict[str, Any]] = [] + + for row in rows_meta[: self.max_rows]: + sid = row.get("system_id") + sidx = row.get("system_index") + sys = None + if sid is not None and str(sid) in sys_by_id: + sys = sys_by_id[str(sid)] + elif sidx is not None: + try: + sys = sys_by_index.get(int(sidx)) + except (TypeError, ValueError): + sys = None + if sys is None: + continue + b = sys.get("bbox") or sys.get("box") or {} + x1, y1, x2, y2 = ( + float(b["x1"]), + float(b["y1"]), + float(b["x2"]), + float(b["y2"]), + ) + # pad crop + pad_x = 0.02 * (x2 - x1) + pad_y = 0.05 * (y2 - y1) + cx1 = max(0, int(x1 - pad_x)) + cy1 = max(0, int(y1 - pad_y)) + cx2 = min(orig_w, int(x2 + pad_x)) + cy2 = min(orig_h, int(y2 + pad_y)) + if cx2 <= cx1 or cy2 <= cy1: + continue + crop = img.crop((cx1, cy1, cx2, cy2)).resize( + (self.row_w, self.row_h), Image.BILINEAR + ) + crop_t = ( + torch.from_numpy(np.array(crop, copy=True).transpose(2, 0, 1)).float() + / 255.0 + ) + # splits in full image → relative to unpadded system x range for heat + # Use crop full-image coords for peak decode consistency + heat = build_l3_x_heat( + row.get("splits") or [], + x_left=float(cx1), + x_right=float(cx2), + heat_len=self.l3_heat_len, + ) + row_images.append(crop_t) + row_heats.append(torch.from_numpy(heat)) + row_meta.append( + { + "system_id": sid, + "x_left": float(cx1), + "x_right": float(cx2), + "y1": float(cy1), + "y2": float(cy2), + "orig_splits": [ + float(sp["x"] if isinstance(sp, dict) else sp) + for sp in (row.get("splits") or []) + ], + "page_box": [ + x1 * sx, + y1 * sy, + x2 * sx, + y2 * sy, + ], + } + ) + + # ensure at least one dummy row for collate stability when L3 empty + if not row_images and "l3" in self.tasks: + row_images.append(torch.zeros(3, self.row_h, self.row_w)) + row_heats.append(torch.zeros(self.l3_heat_len)) + row_meta.append( + { + "system_id": None, + "x_left": 0.0, + "x_right": 1.0, + "y1": 0.0, + "y2": 1.0, + "orig_splits": [], + "page_box": [0, 0, 1, 1], + "dummy": True, + } + ) + + return { + "id": layout.get("id") or sample_dir.name, + "path": str(sample_dir), + "page": page_t, + "orig_size": (orig_h, orig_w), + "l2_heat": torch.from_numpy(l2_heat), + "l2_boxes": torch.tensor(l2_boxes, dtype=torch.float32) + if l2_boxes + else torch.zeros(0, 4), + "row_images": row_images, + "row_heats": row_heats, + "row_meta": row_meta, + "layout": layout, + } + + +def collate_layout(batch: list[dict[str, Any]]) -> dict[str, Any]: + """Collate variable-length rows by flattening L3 crops across the batch.""" + pages = torch.stack([b["page"] for b in batch], dim=0) + l2_heats = torch.stack([b["l2_heat"] for b in batch], dim=0) + ids = [b["id"] for b in batch] + l2_boxes = [b["l2_boxes"] for b in batch] + + row_images: list[torch.Tensor] = [] + row_heats: list[torch.Tensor] = [] + row_meta: list[dict[str, Any]] = [] + row_batch_idx: list[int] = [] + for bi, b in enumerate(batch): + for img, heat, meta in zip(b["row_images"], b["row_heats"], b["row_meta"]): + if meta.get("dummy"): + continue + row_images.append(img) + row_heats.append(heat) + row_meta.append(meta) + row_batch_idx.append(bi) + + out: dict[str, Any] = { + "ids": ids, + "page": pages, + "l2_heat": l2_heats, + "l2_boxes": l2_boxes, + "paths": [b["path"] for b in batch], + "orig_sizes": [b["orig_size"] for b in batch], + } + if row_images: + out["row_images"] = torch.stack(row_images, dim=0) + out["row_heats"] = torch.stack(row_heats, dim=0) + out["row_meta"] = row_meta + out["row_batch_idx"] = torch.tensor(row_batch_idx, dtype=torch.long) + else: + # empty L3 — create zero batch for model path + h = batch[0]["row_images"][0].shape[-2] if batch[0]["row_images"] else 64 + w = batch[0]["row_images"][0].shape[-1] if batch[0]["row_images"] else 256 + hl = batch[0]["row_heats"][0].shape[0] if batch[0]["row_heats"] else 128 + out["row_images"] = torch.zeros(0, 3, h, w) + out["row_heats"] = torch.zeros(0, hl) + out["row_meta"] = [] + out["row_batch_idx"] = torch.zeros(0, dtype=torch.long) + return out diff --git a/train/enpu_train/data/synthetic.py b/train/enpu_train/data/synthetic.py new file mode 100644 index 0000000..264766f --- /dev/null +++ b/train/enpu_train/data/synthetic.py @@ -0,0 +1,185 @@ +"""Procedural layout samples for toy training when real data is scarce (#95).""" + +from __future__ import annotations + +import json +import random +from pathlib import Path +from typing import Any + +import numpy as np +from PIL import Image, ImageDraw + + +def make_synthetic_layout_sample( + out_dir: str | Path, + *, + sample_id: str, + width: int = 640, + height: int = 800, + n_systems: int | None = None, + measures_per_row: int | None = None, + seed: int | None = None, +) -> dict[str, Any]: + """Draw a simple jianpu-like page and write layout.json + image.png. + + Geometry is exact so it can act as GT for L2 boxes and L3 interior splits. + """ + rng = random.Random(seed) + out_dir = Path(out_dir) + out_dir.mkdir(parents=True, exist_ok=True) + + n_systems = n_systems if n_systems is not None else rng.randint(3, 6) + measures_per_row = ( + measures_per_row if measures_per_row is not None else rng.randint(3, 5) + ) + n_splits = measures_per_row - 1 + + img = Image.new("RGB", (width, height), (255, 255, 255)) + draw = ImageDraw.Draw(img) + + # Title band + title_box = { + "x1": width * 0.2, + "y1": height * 0.04, + "x2": width * 0.8, + "y2": height * 0.09, + } + draw.rectangle( + [title_box["x1"], title_box["y1"], title_box["x2"], title_box["y2"]], + fill=(30, 30, 30), + ) + + score_y0 = height * 0.12 + score_y1 = height * 0.92 + score_region = {"x1": 0.0, "y1": score_y0, "x2": float(width), "y2": score_y1} + + margin_x = width * 0.06 + usable_w = width - 2 * margin_x + band_h = min(70.0, (score_y1 - score_y0) / (n_systems + 1) * 0.7) + gap = (score_y1 - score_y0 - n_systems * band_h) / (n_systems + 1) + + systems: list[dict[str, Any]] = [] + rows: list[dict[str, Any]] = [] + + y = score_y0 + gap + for si in range(n_systems): + y1 = y + y2 = y + band_h + x1, x2 = margin_x, width - margin_x + # staff ink + draw.rectangle([x1, y1, x2, y2], fill=(40, 40, 40)) + # digit-like blobs + for k in range(measures_per_row * 3): + cx = x1 + (k + 0.5) * usable_w / (measures_per_row * 3) + draw.rectangle( + [cx - 4, y1 + band_h * 0.25, cx + 4, y1 + band_h * 0.75], + fill=(10, 10, 10), + ) + + # interior split lines (white gaps as "barlines") + splits = [] + for j in range(n_splits): + # equal measure widths + sx = x1 + (j + 1) * usable_w / measures_per_row + draw.line([(sx, y1), (sx, y2)], fill=(255, 255, 255), width=3) + splits.append( + { + "id": f"s{si}-{j}", + "x": float(sx), + "y1": float(y1), + "y2": float(y2), + "source": "synth", + } + ) + + edges = [x1] + [s["x"] for s in splits] + [x2] + measures = [] + for mi in range(len(edges) - 1): + measures.append( + { + "id": f"l3-m{si}-{mi}", + "label": f"m{mi + 1}", + "box": { + "x1": float(edges[mi]), + "y1": float(y1), + "x2": float(edges[mi + 1]), + "y2": float(y2), + }, + } + ) + + sid = f"l2-sys{si}" + systems.append( + { + "id": sid, + "index": si, + "bbox": { + "x1": float(x1), + "y1": float(y1), + "x2": float(x2), + "y2": float(y2), + }, + "kind": "system", + } + ) + rows.append( + { + "system_id": sid, + "system_index": si, + "splits": splits, + "measures": measures, + } + ) + y = y2 + gap + + layout: dict[str, Any] = { + "layout_schema_version": "0.1", + "kind": "enpu-layout-gt", + "id": sample_id, + "image": {"path": "image.png", "width": width, "height": height}, + "meta": {"title": sample_id, "source": "synthetic"}, + "l1": { + "score_region": score_region, + "title": title_box, + "regions": [ + {"role": "title", "box": title_box}, + {"role": "score", "box": score_region}, + ], + }, + "l2": {"systems": systems}, + "l3": {"rows": rows}, + "source": {"type": "synthetic", "seed": seed}, + } + + img_path = out_dir / "image.png" + img.save(img_path) + (out_dir / "layout.json").write_text( + json.dumps(layout, ensure_ascii=False, indent=2) + "\n", + encoding="utf-8", + ) + return layout + + +def generate_synthetic_set( + root: str | Path, + n: int = 8, + *, + seed: int = 0, +) -> list[Path]: + """Write ``root/S001`` … samples; return directories.""" + root = Path(root) + paths: list[Path] = [] + rng = random.Random(seed) + for i in range(n): + sid = f"S{i + 1:03d}_synth" + d = root / sid + make_synthetic_layout_sample( + d, + sample_id=sid, + width=rng.choice([512, 640, 720]), + height=rng.choice([640, 800, 900]), + seed=seed + i * 17, + ) + paths.append(d) + return paths diff --git a/train/enpu_train/engine/__init__.py b/train/enpu_train/engine/__init__.py new file mode 100644 index 0000000..91deb58 --- /dev/null +++ b/train/enpu_train/engine/__init__.py @@ -0,0 +1,3 @@ +from enpu_train.engine.trainer import TrainConfig, evaluate, train_loop + +__all__ = ["TrainConfig", "evaluate", "train_loop"] diff --git a/train/enpu_train/engine/trainer.py b/train/enpu_train/engine/trainer.py new file mode 100644 index 0000000..00ca9ba --- /dev/null +++ b/train/enpu_train/engine/trainer.py @@ -0,0 +1,205 @@ +"""Train / eval loops for LayoutNet (#95).""" + +from __future__ import annotations + +import json +import time +from dataclasses import asdict, dataclass, field +from pathlib import Path +from typing import Any + +import torch +from torch.utils.data import DataLoader + +from enpu_train.losses.heat import heatmap_bce_loss +from enpu_train.metrics.layout_metrics import evaluate_batch +from enpu_train.models.layout_net import LayoutNet, LayoutNetConfig + + +@dataclass +class TrainConfig: + tasks: list[str] = field(default_factory=lambda: ["l2", "l3"]) + epochs: int = 3 + batch_size: int = 2 + lr: float = 1e-3 + weight_decay: float = 1e-4 + l2_loss_weight: float = 1.0 + l3_loss_weight: float = 1.5 + device: str = "cpu" + page_h: int = 384 + page_w: int = 512 + row_h: int = 64 + row_w: int = 256 + l2_heat_len: int = 128 + l3_heat_len: int = 128 + num_workers: int = 0 + out_dir: str = "runs/toy" + log_every: int = 1 + + +def _move_batch(batch: dict[str, Any], device: torch.device) -> dict[str, Any]: + out = dict(batch) + for k in ("page", "l2_heat", "row_images", "row_heats", "row_batch_idx"): + if k in out and torch.is_tensor(out[k]): + out[k] = out[k].to(device) + return out + + +def compute_loss( + model: LayoutNet, + batch: dict[str, Any], + cfg: TrainConfig, +) -> tuple[torch.Tensor, dict[str, float]]: + out = model(page=batch["page"], rows=batch["row_images"]) + loss = batch["page"].new_zeros(()) + parts: dict[str, float] = {} + + if "l2" in cfg.tasks and "l2_logits" in out: + l2 = heatmap_bce_loss(out["l2_logits"], batch["l2_heat"]) + loss = loss + cfg.l2_loss_weight * l2 + parts["l2"] = float(l2.detach().cpu()) + + if "l3" in cfg.tasks and "l3_logits" in out and batch["row_heats"].shape[0] > 0: + l3 = heatmap_bce_loss(out["l3_logits"], batch["row_heats"]) + loss = loss + cfg.l3_loss_weight * l3 + parts["l3"] = float(l3.detach().cpu()) + elif "l3" in cfg.tasks: + parts["l3"] = 0.0 + + parts["total"] = float(loss.detach().cpu()) + return loss, parts + + +@torch.no_grad() +def evaluate( + model: LayoutNet, + loader: DataLoader, + cfg: TrainConfig, +) -> dict[str, float]: + model.eval() + device = torch.device(cfg.device) + acc: dict[str, list[float]] = { + "l2_mean_iou": [], + "l3_split_count_mae": [], + "l3_split_count_exact": [], + "l3_mean_abs_x_error": [], + "loss": [], + } + for batch in loader: + batch = _move_batch(batch, device) + loss, _ = compute_loss(model, batch, cfg) + acc["loss"].append(float(loss.cpu())) + out = model(page=batch["page"], rows=batch["row_images"]) + m = evaluate_batch( + out, batch, page_h=float(cfg.page_h), page_w=float(cfg.page_w) + ) + for k in ( + "l2_mean_iou", + "l3_split_count_mae", + "l3_split_count_exact", + "l3_mean_abs_x_error", + ): + v = m.get(k) + if v == v: # not nan + acc[k].append(float(v)) + + def mean(xs: list[float]) -> float: + return float(sum(xs) / len(xs)) if xs else float("nan") + + return {k: mean(v) for k, v in acc.items()} + + +def train_loop( + model: LayoutNet, + train_loader: DataLoader, + val_loader: DataLoader | None, + cfg: TrainConfig, +) -> dict[str, Any]: + device = torch.device(cfg.device) + model = model.to(device) + opt = torch.optim.AdamW( + model.parameters(), lr=cfg.lr, weight_decay=cfg.weight_decay + ) + out_dir = Path(cfg.out_dir) + out_dir.mkdir(parents=True, exist_ok=True) + + history: list[dict[str, Any]] = [] + best_val = float("inf") + best_path = out_dir / "best.pt" + + t0 = time.time() + for epoch in range(1, cfg.epochs + 1): + model.train() + epoch_losses: list[float] = [] + for step, batch in enumerate(train_loader, start=1): + batch = _move_batch(batch, device) + opt.zero_grad(set_to_none=True) + loss, parts = compute_loss(model, batch, cfg) + loss.backward() + opt.step() + epoch_losses.append(parts["total"]) + if step % max(1, cfg.log_every) == 0: + print( + f"epoch {epoch} step {step} " + f"loss={parts['total']:.4f} " + f"l2={parts.get('l2', float('nan')):.4f} " + f"l3={parts.get('l3', float('nan')):.4f}" + ) + + record: dict[str, Any] = { + "epoch": epoch, + "train_loss": float(sum(epoch_losses) / max(1, len(epoch_losses))), + } + if val_loader is not None: + metrics = evaluate(model, val_loader, cfg) + record["val"] = metrics + print( + f"epoch {epoch} val loss={metrics['loss']:.4f} " + f"l2_iou={metrics['l2_mean_iou']:.4f} " + f"l3_x_err={metrics['l3_mean_abs_x_error']:.4f} " + f"l3_count_mae={metrics['l3_split_count_mae']:.4f}" + ) + score = metrics["loss"] + if score < best_val: + best_val = score + torch.save( + { + "model": model.state_dict(), + "cfg": asdict(cfg), + "layout_net": asdict(model.cfg) + if hasattr(model.cfg, "__dataclass_fields__") + else {}, + "epoch": epoch, + "metrics": metrics, + }, + best_path, + ) + # always save last + torch.save( + { + "model": model.state_dict(), + "cfg": asdict(cfg), + "epoch": epoch, + }, + out_dir / "last.pt", + ) + history.append(record) + + (out_dir / "history.json").write_text( + json.dumps(history, indent=2) + "\n", encoding="utf-8" + ) + return { + "history": history, + "best_path": str(best_path) if best_path.is_file() else str(out_dir / "last.pt"), + "seconds": time.time() - t0, + } + + +def build_model_from_cfg(cfg: TrainConfig) -> LayoutNet: + return LayoutNet( + LayoutNetConfig( + l2_heat_len=cfg.l2_heat_len, + l3_heat_len=cfg.l3_heat_len, + tasks=tuple(cfg.tasks), + ) + ) diff --git a/train/enpu_train/export/__init__.py b/train/enpu_train/export/__init__.py new file mode 100644 index 0000000..eb7812d --- /dev/null +++ b/train/enpu_train/export/__init__.py @@ -0,0 +1,3 @@ +from enpu_train.export.weights import export_onnx, export_state_dict + +__all__ = ["export_onnx", "export_state_dict"] diff --git a/train/enpu_train/export/weights.py b/train/enpu_train/export/weights.py new file mode 100644 index 0000000..c6ac147 --- /dev/null +++ b/train/enpu_train/export/weights.py @@ -0,0 +1,122 @@ +"""Export trained weights for later core integration (#95).""" + +from __future__ import annotations + +from pathlib import Path +from typing import Any + +import torch + +from enpu_train.models.layout_net import LayoutNet, LayoutNetConfig + + +def load_checkpoint(path: str | Path, device: str = "cpu") -> tuple[LayoutNet, dict]: + ckpt = torch.load(path, map_location=device, weights_only=False) + cfg_d = ckpt.get("cfg") or {} + net_cfg = LayoutNetConfig( + l2_heat_len=int(cfg_d.get("l2_heat_len", 128)), + l3_heat_len=int(cfg_d.get("l3_heat_len", 128)), + tasks=tuple(cfg_d.get("tasks") or ("l2", "l3")), + ) + model = LayoutNet(net_cfg) + model.load_state_dict(ckpt["model"]) + model.to(device) + model.eval() + return model, ckpt + + +def export_state_dict(ckpt_path: str | Path, out_path: str | Path) -> Path: + """Copy/slim to a portable state_dict file.""" + out_path = Path(out_path) + out_path.parent.mkdir(parents=True, exist_ok=True) + model, ckpt = load_checkpoint(ckpt_path) + payload: dict[str, Any] = { + "format": "enpu_layout_net_v0", + "model": model.state_dict(), + "tasks": list(model.cfg.tasks), + "l2_heat_len": model.cfg.l2_heat_len, + "l3_heat_len": model.cfg.l3_heat_len, + "source_ckpt": str(ckpt_path), + "train_metrics": ckpt.get("metrics"), + "note": ( + "Load in core via future engine=learned_l1l3 (P2). " + "Map L2 y-peaks → system bands; L3 x-peaks → barlines/splits; " + "then normalize_splits + splits_to_measures." + ), + } + torch.save(payload, out_path) + return out_path + + +def export_onnx( + ckpt_path: str | Path, + out_dir: str | Path, + *, + page_h: int = 384, + page_w: int = 512, + row_h: int = 64, + row_w: int = 256, + opset: int = 17, +) -> dict[str, str]: + """Export separate ONNX graphs for L2 page head and L3 row head.""" + out_dir = Path(out_dir) + out_dir.mkdir(parents=True, exist_ok=True) + model, _ = load_checkpoint(ckpt_path) + paths: dict[str, str] = {} + + # L2 + if "l2" in model.cfg.tasks: + class _L2(torch.nn.Module): + def __init__(self, m: LayoutNet) -> None: + super().__init__() + self.m = m + + def forward(self, page: torch.Tensor) -> torch.Tensor: + return self.m.l2(page) + + l2 = _L2(model).eval() + dummy = torch.randn(1, 3, page_h, page_w) + p = out_dir / "l2_page.onnx" + torch.onnx.export( + l2, + dummy, + str(p), + input_names=["page"], + output_names=["l2_logits"], + dynamic_axes={"page": {0: "batch"}, "l2_logits": {0: "batch"}}, + opset_version=opset, + ) + paths["l2"] = str(p) + + # L3 + if "l3" in model.cfg.tasks: + class _L3(torch.nn.Module): + def __init__(self, m: LayoutNet) -> None: + super().__init__() + self.m = m + + def forward(self, rows: torch.Tensor) -> torch.Tensor: + return self.m.l3(rows) + + l3 = _L3(model).eval() + dummy = torch.randn(1, 3, row_h, row_w) + p = out_dir / "l3_row.onnx" + torch.onnx.export( + l3, + dummy, + str(p), + input_names=["row"], + output_names=["l3_logits"], + dynamic_axes={"row": {0: "batch"}, "l3_logits": {0: "batch"}}, + opset_version=opset, + ) + paths["l3"] = str(p) + + (out_dir / "README_export.txt").write_text( + "EnPu layout ONNX export (#95)\n" + "l2_page.onnx: input page NCHW float0-1 → l2_logits (N, heat_len)\n" + "l3_row.onnx: input row crop NCHW → l3_logits (N, heat_len)\n" + "Post-process: peak decode → structure.barlines / L2 items; see train/README.md\n", + encoding="utf-8", + ) + return paths diff --git a/train/enpu_train/losses/__init__.py b/train/enpu_train/losses/__init__.py new file mode 100644 index 0000000..1f3906c --- /dev/null +++ b/train/enpu_train/losses/__init__.py @@ -0,0 +1,3 @@ +from enpu_train.losses.heat import heatmap_bce_loss + +__all__ = ["heatmap_bce_loss"] diff --git a/train/enpu_train/losses/heat.py b/train/enpu_train/losses/heat.py new file mode 100644 index 0000000..d9a934e --- /dev/null +++ b/train/enpu_train/losses/heat.py @@ -0,0 +1,20 @@ +"""Heatmap losses for L2/L3 1D heads.""" + +from __future__ import annotations + +import torch +import torch.nn.functional as F + + +def heatmap_bce_loss( + logits: torch.Tensor, + target: torch.Tensor, + *, + pos_weight: float = 3.0, +) -> torch.Tensor: + """BCE with logits; upweight positives (sparse peaks).""" + if logits.numel() == 0: + return logits.sum() * 0.0 + # target in [0,1] + pw = torch.tensor(pos_weight, device=logits.device, dtype=logits.dtype) + return F.binary_cross_entropy_with_logits(logits, target, pos_weight=pw) diff --git a/train/enpu_train/metrics/__init__.py b/train/enpu_train/metrics/__init__.py new file mode 100644 index 0000000..b1213fd --- /dev/null +++ b/train/enpu_train/metrics/__init__.py @@ -0,0 +1,13 @@ +from enpu_train.metrics.layout_metrics import ( + boxes_mean_iou, + evaluate_batch, + l2_peak_boxes_iou, + split_x_metrics, +) + +__all__ = [ + "boxes_mean_iou", + "evaluate_batch", + "l2_peak_boxes_iou", + "split_x_metrics", +] diff --git a/train/enpu_train/metrics/layout_metrics.py b/train/enpu_train/metrics/layout_metrics.py new file mode 100644 index 0000000..a3b58cd --- /dev/null +++ b/train/enpu_train/metrics/layout_metrics.py @@ -0,0 +1,199 @@ +"""Hard metrics: L2 box IoU (approx) + L3 split count / mean abs x (#94 / #86).""" + +from __future__ import annotations + +from typing import Any + +import numpy as np +import torch + +from enpu_train.data.dataset import decode_peaks + + +def split_x_metrics( + gt_xs: list[float], + pred_xs: list[float], + *, + max_dist: float = 12.0, +) -> dict[str, float]: + """Match split x positions (same idea as core barline_x_metrics).""" + gt = sorted(float(x) for x in gt_xs) + pred = sorted(float(x) for x in pred_xs) + if not gt and not pred: + return { + "split_count_mae": 0.0, + "split_count_exact": 1.0, + "split_mean_abs_x_error": 0.0, + "n_gt": 0.0, + "n_pred": 0.0, + } + used: set[int] = set() + dists: list[float] = [] + tp = 0 + for gx in gt: + best_i, best_d = -1, max_dist + 1 + for i, px in enumerate(pred): + if i in used: + continue + d = abs(px - gx) + if d < best_d: + best_d, best_i = d, i + if best_i >= 0 and best_d <= max_dist: + used.add(best_i) + tp += 1 + dists.append(best_d) + return { + "split_count_mae": float(abs(len(pred) - len(gt))), + "split_count_exact": 1.0 if len(pred) == len(gt) else 0.0, + "split_mean_abs_x_error": float(sum(dists) / len(dists)) if dists else float("nan"), + "n_gt": float(len(gt)), + "n_pred": float(len(pred)), + "tp": float(tp), + "fp": float(len(pred) - tp), + "fn": float(len(gt) - tp), + } + + +def box_iou(a: list[float] | np.ndarray, b: list[float] | np.ndarray) -> float: + ax1, ay1, ax2, ay2 = map(float, a) + bx1, by1, bx2, by2 = map(float, b) + ix1, iy1 = max(ax1, bx1), max(ay1, by1) + ix2, iy2 = min(ax2, bx2), min(ay2, by2) + iw, ih = max(0.0, ix2 - ix1), max(0.0, iy2 - iy1) + inter = iw * ih + if inter <= 0: + return 0.0 + area_a = max(0.0, ax2 - ax1) * max(0.0, ay2 - ay1) + area_b = max(0.0, bx2 - bx1) * max(0.0, by2 - by1) + union = area_a + area_b - inter + return float(inter / union) if union > 0 else 0.0 + + +def boxes_mean_iou( + gt_boxes: list[list[float]], + pred_boxes: list[list[float]], +) -> float: + """Greedy match by IoU; mean over GT (unmatched = 0).""" + if not gt_boxes: + return 1.0 if not pred_boxes else 0.0 + remaining = list(range(len(pred_boxes))) + ious: list[float] = [] + for g in gt_boxes: + best_j, best = -1, 0.0 + for j in remaining: + v = box_iou(g, pred_boxes[j]) + if v > best: + best, best_j = v, j + if best_j >= 0 and best > 0: + remaining.remove(best_j) + ious.append(best) + else: + ious.append(0.0) + return float(sum(ious) / len(ious)) + + +def l2_peaks_to_bands( + heat: np.ndarray, + *, + page_h: float, + page_w: float, + band_frac: float = 0.08, +) -> list[list[float]]: + """Decode y-peaks to full-width horizontal bands (approx L2 boxes).""" + peaks = decode_peaks(heat, min_prominence=0.25, min_gap=max(3, len(heat) // 40)) + boxes = [] + half = band_frac * page_h * 0.5 + for p in peaks: + cy = p / max(1, len(heat) - 1) * page_h + boxes.append([0.0, max(0.0, cy - half), page_w, min(page_h, cy + half)]) + return boxes + + +def l2_peak_boxes_iou( + heat: np.ndarray, + gt_boxes: torch.Tensor | list, + *, + page_h: float, + page_w: float, +) -> float: + if isinstance(gt_boxes, torch.Tensor): + gt = gt_boxes.detach().cpu().tolist() + else: + gt = list(gt_boxes) + pred = l2_peaks_to_bands(heat, page_h=page_h, page_w=page_w) + return boxes_mean_iou(gt, pred) + + +def l3_heat_to_xs( + heat: np.ndarray, + *, + x_left: float, + x_right: float, + min_prominence: float = 0.3, +) -> list[float]: + peaks = decode_peaks( + heat, + min_prominence=min_prominence, + min_gap=max(3, len(heat) // 32), + ) + width = max(1e-3, x_right - x_left) + xs = [] + for p in peaks: + rel = p / max(1, len(heat) - 1) + xs.append(x_left + rel * width) + return xs + + +def evaluate_batch( + model_out: dict[str, torch.Tensor], + batch: dict[str, Any], + *, + page_h: float, + page_w: float, +) -> dict[str, float]: + """Aggregate L2 IoU and L3 split metrics over a batch.""" + stats = { + "l2_mean_iou": [], + "l3_split_count_mae": [], + "l3_split_count_exact": [], + "l3_mean_abs_x_error": [], + } + + if "l2_logits" in model_out: + probs = torch.sigmoid(model_out["l2_logits"]).detach().cpu().numpy() + for i in range(probs.shape[0]): + gt = batch["l2_boxes"][i] + iou = l2_peak_boxes_iou( + probs[i], gt, page_h=page_h, page_w=page_w + ) + stats["l2_mean_iou"].append(iou) + + if "l3_logits" in model_out and batch["row_heats"].numel() > 0: + probs = torch.sigmoid(model_out["l3_logits"]).detach().cpu().numpy() + for i, meta in enumerate(batch["row_meta"]): + pred_xs = l3_heat_to_xs( + probs[i], + x_left=meta["x_left"], + x_right=meta["x_right"], + ) + m = split_x_metrics( + meta.get("orig_splits") or [], + pred_xs, + max_dist=max(12.0, 0.02 * (meta["x_right"] - meta["x_left"])), + ) + stats["l3_split_count_mae"].append(m["split_count_mae"]) + stats["l3_split_count_exact"].append(m["split_count_exact"]) + if m["split_mean_abs_x_error"] == m["split_mean_abs_x_error"]: # not nan + stats["l3_mean_abs_x_error"].append(m["split_mean_abs_x_error"]) + + def _mean(xs: list[float]) -> float: + return float(sum(xs) / len(xs)) if xs else float("nan") + + return { + "l2_mean_iou": _mean(stats["l2_mean_iou"]), + "l3_split_count_mae": _mean(stats["l3_split_count_mae"]), + "l3_split_count_exact": _mean(stats["l3_split_count_exact"]), + "l3_mean_abs_x_error": _mean(stats["l3_mean_abs_x_error"]), + "n_l2": float(len(stats["l2_mean_iou"])), + "n_l3_rows": float(len(stats["l3_split_count_mae"])), + } diff --git a/train/enpu_train/models/__init__.py b/train/enpu_train/models/__init__.py new file mode 100644 index 0000000..3aa8976 --- /dev/null +++ b/train/enpu_train/models/__init__.py @@ -0,0 +1,3 @@ +from enpu_train.models.layout_net import LayoutNet, LayoutNetConfig + +__all__ = ["LayoutNet", "LayoutNetConfig"] diff --git a/train/enpu_train/models/layout_net.py b/train/enpu_train/models/layout_net.py new file mode 100644 index 0000000..0cadd91 --- /dev/null +++ b/train/enpu_train/models/layout_net.py @@ -0,0 +1,122 @@ +"""Lightweight L2 (page y-heat) + L3 (row x-heat) networks (#95 / #94 path A).""" + +from __future__ import annotations + +from dataclasses import dataclass + +import torch +import torch.nn as nn +import torch.nn.functional as F + + +class ConvBNReLU(nn.Module): + def __init__(self, c_in: int, c_out: int, k: int = 3, s: int = 1) -> None: + super().__init__() + p = k // 2 + self.net = nn.Sequential( + nn.Conv2d(c_in, c_out, k, stride=s, padding=p, bias=False), + nn.BatchNorm2d(c_out), + nn.ReLU(inplace=True), + ) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return self.net(x) + + +class PageL2Head(nn.Module): + """Page image → 1D heatmap along height (system row centers).""" + + def __init__(self, heat_len: int = 128, base: int = 16) -> None: + super().__init__() + self.heat_len = heat_len + self.enc = nn.Sequential( + ConvBNReLU(3, base, 3, 2), + ConvBNReLU(base, base * 2, 3, 2), + ConvBNReLU(base * 2, base * 4, 3, 2), + ConvBNReLU(base * 4, base * 4, 3, 2), + ) + self.proj = nn.Sequential( + nn.AdaptiveAvgPool2d((heat_len, 1)), # (B,C,H',1) roughly then we force + ) + # After pool we may not get exact heat_len; use linear on flattened H + self.head = nn.Sequential( + nn.Conv1d(base * 4, base * 2, 3, padding=1), + nn.ReLU(inplace=True), + nn.Conv1d(base * 2, 1, 1), + ) + + def forward(self, page: torch.Tensor) -> torch.Tensor: + # page: B,3,H,W + f = self.enc(page) # B,C,h,w + # collapse width + f = f.mean(dim=3) # B,C,h + # interpolate along H to heat_len + f = F.interpolate( + f.unsqueeze(-1), + size=(self.heat_len, 1), + mode="bilinear", + align_corners=False, + ).squeeze(-1) # B,C,heat_len + logits = self.head(f).squeeze(1) # B, heat_len + return logits + + +class RowL3Head(nn.Module): + """Row crop → 1D heatmap along width (interior splits).""" + + def __init__(self, heat_len: int = 128, base: int = 16) -> None: + super().__init__() + self.heat_len = heat_len + self.enc = nn.Sequential( + ConvBNReLU(3, base, 3, 2), + ConvBNReLU(base, base * 2, 3, 2), + ConvBNReLU(base * 2, base * 4, 3, 2), + ConvBNReLU(base * 4, base * 4, 3, 2), + ) + self.head = nn.Sequential( + nn.Conv1d(base * 4, base * 2, 3, padding=1), + nn.ReLU(inplace=True), + nn.Conv1d(base * 2, 1, 1), + ) + + def forward(self, rows: torch.Tensor) -> torch.Tensor: + if rows.numel() == 0: + return rows.new_zeros((0, self.heat_len)) + f = self.enc(rows) # B,C,h,w + f = f.mean(dim=2) # B,C,w + f = F.interpolate( + f.unsqueeze(2), + size=(1, self.heat_len), + mode="bilinear", + align_corners=False, + ).squeeze(2) # B,C,heat_len + logits = self.head(f).squeeze(1) + return logits + + +@dataclass +class LayoutNetConfig: + l2_heat_len: int = 128 + l3_heat_len: int = 128 + base_channels: int = 16 + tasks: tuple[str, ...] = ("l2", "l3") + + +class LayoutNet(nn.Module): + def __init__(self, cfg: LayoutNetConfig | None = None) -> None: + super().__init__() + self.cfg = cfg or LayoutNetConfig() + self.l2 = PageL2Head(self.cfg.l2_heat_len, self.cfg.base_channels) + self.l3 = RowL3Head(self.cfg.l3_heat_len, self.cfg.base_channels) + + def forward( + self, + page: torch.Tensor | None = None, + rows: torch.Tensor | None = None, + ) -> dict[str, torch.Tensor]: + out: dict[str, torch.Tensor] = {} + if page is not None and "l2" in self.cfg.tasks: + out["l2_logits"] = self.l2(page) + if rows is not None and "l3" in self.cfg.tasks: + out["l3_logits"] = self.l3(rows) + return out diff --git a/train/enpu_train/viz.py b/train/enpu_train/viz.py new file mode 100644 index 0000000..dcea527 --- /dev/null +++ b/train/enpu_train/viz.py @@ -0,0 +1,60 @@ +"""Visualize layout GT / predictions: page boxes + L3 vertical lines.""" + +from __future__ import annotations + +from pathlib import Path +from typing import Any + +import numpy as np +from PIL import Image, ImageDraw + + +def draw_layout_overlay( + image: Image.Image | np.ndarray | str | Path, + layout: dict[str, Any], + *, + out_path: str | Path | None = None, +) -> Image.Image: + if isinstance(image, (str, Path)): + img = Image.open(image).convert("RGB") + elif isinstance(image, np.ndarray): + img = Image.fromarray(image.astype(np.uint8)).convert("RGB") + else: + img = image.convert("RGB") + draw = ImageDraw.Draw(img) + + l1 = layout.get("l1") or {} + for key, color in ( + ("score_region", (0, 180, 0)), + ("title", (0, 0, 220)), + ("key_time", (180, 0, 180)), + ): + box = l1.get(key) + if box: + draw.rectangle( + [box["x1"], box["y1"], box["x2"], box["y2"]], + outline=color, + width=3, + ) + + for s in (layout.get("l2") or {}).get("systems") or []: + b = s.get("bbox") or s.get("box") or {} + draw.rectangle( + [b["x1"], b["y1"], b["x2"], b["y2"]], + outline=(220, 120, 0), + width=2, + ) + + for row in (layout.get("l3") or {}).get("rows") or []: + for sp in row.get("splits") or []: + x = float(sp["x"] if isinstance(sp, dict) else sp) + y1 = float(sp.get("y1", 0) if isinstance(sp, dict) else 0) + y2 = float(sp.get("y2", img.height) if isinstance(sp, dict) else img.height) + if y2 <= y1: + y1, y2 = 0, img.height + draw.line([(x, y1), (x, y2)], fill=(220, 30, 30), width=2) + + if out_path is not None: + Path(out_path).parent.mkdir(parents=True, exist_ok=True) + img.save(out_path) + return img diff --git a/train/requirements.txt b/train/requirements.txt new file mode 100644 index 0000000..58185b0 --- /dev/null +++ b/train/requirements.txt @@ -0,0 +1,6 @@ +# EnPu train framework (#95) — install in a venv separate from core if preferred +torch>=2.0 +numpy>=1.24 +Pillow>=10.0 +PyYAML>=6.0 +tqdm>=4.65 diff --git a/train/scripts/eval.py b/train/scripts/eval.py new file mode 100644 index 0000000..6b309e5 --- /dev/null +++ b/train/scripts/eval.py @@ -0,0 +1,65 @@ +#!/usr/bin/env python3 +"""Evaluate a checkpoint on layout samples — L2 IoU + L3 split x metrics (#95).""" + +from __future__ import annotations + +import argparse +import json +import sys +from pathlib import Path + +import torch +import yaml +from torch.utils.data import DataLoader + +ROOT = Path(__file__).resolve().parents[1] +if str(ROOT) not in sys.path: + sys.path.insert(0, str(ROOT)) + +from enpu_train.data.dataset import LayoutDataset, collate_layout +from enpu_train.engine.trainer import TrainConfig, evaluate +from enpu_train.export.weights import load_checkpoint + + +def main(argv: list[str] | None = None) -> int: + ap = argparse.ArgumentParser() + ap.add_argument("--ckpt", type=Path, required=True) + ap.add_argument("--data", type=Path, default=ROOT.parent / "samples" / "layout") + ap.add_argument("--device", type=str, default="cpu") + ap.add_argument("--out", type=Path, default=None) + args = ap.parse_args(argv) + + model, ckpt = load_checkpoint(args.ckpt, device=args.device) + cfg_d = ckpt.get("cfg") or {} + tcfg = TrainConfig( + tasks=list(cfg_d.get("tasks") or ["l2", "l3"]), + device=args.device, + page_h=int(cfg_d.get("page_h", 384)), + page_w=int(cfg_d.get("page_w", 512)), + row_h=int(cfg_d.get("row_h", 64)), + row_w=int(cfg_d.get("row_w", 256)), + l2_heat_len=int(cfg_d.get("l2_heat_len", 128)), + l3_heat_len=int(cfg_d.get("l3_heat_len", 128)), + batch_size=1, + ) + # rebuild model is already loaded; ensure cfg match + ds = LayoutDataset( + args.data, + page_size=(tcfg.page_h, tcfg.page_w), + row_size=(tcfg.row_h, tcfg.row_w), + l2_heat_len=tcfg.l2_heat_len, + l3_heat_len=tcfg.l3_heat_len, + tasks=tuple(tcfg.tasks), + augment=False, + ) + loader = DataLoader(ds, batch_size=1, shuffle=False, collate_fn=collate_layout) + metrics = evaluate(model, loader, tcfg) + print(json.dumps(metrics, indent=2)) + if args.out: + args.out.parent.mkdir(parents=True, exist_ok=True) + args.out.write_text(json.dumps(metrics, indent=2) + "\n", encoding="utf-8") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/train/scripts/export_from_enpu_project.py b/train/scripts/export_from_enpu_project.py new file mode 100644 index 0000000..5cca147 --- /dev/null +++ b/train/scripts/export_from_enpu_project.py @@ -0,0 +1,34 @@ +#!/usr/bin/env python3 +"""Thin wrapper: export .enpu.json → layout sample via core layout_gt (#93/#95).""" + +from __future__ import annotations + +import argparse +import sys +from pathlib import Path + +REPO = Path(__file__).resolve().parents[2] +CORE = REPO / "core" +if str(CORE) not in sys.path: + sys.path.insert(0, str(CORE)) + +from app.layout_gt.export import export_project_to_sample_dir # noqa: E402 + + +def main() -> int: + ap = argparse.ArgumentParser() + ap.add_argument("--project", "-p", type=Path, required=True) + ap.add_argument("--out", "-o", type=Path, required=True) + ap.add_argument("--sample-id", type=str, default=None) + args = ap.parse_args() + sample = export_project_to_sample_dir( + args.project, + args.out, + sample_id=args.sample_id, + ) + print("exported", args.out / "layout.json", "id=", sample.get("id")) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/train/scripts/train.py b/train/scripts/train.py new file mode 100644 index 0000000..81019e4 --- /dev/null +++ b/train/scripts/train.py @@ -0,0 +1,178 @@ +#!/usr/bin/env python3 +"""Train L2+L3 layout model (toy MVP) — #95.""" + +from __future__ import annotations + +import argparse +import random +import sys +from pathlib import Path + +import torch +import yaml +from torch.utils.data import DataLoader, Subset + +ROOT = Path(__file__).resolve().parents[1] +if str(ROOT) not in sys.path: + sys.path.insert(0, str(ROOT)) + +from enpu_train.data.dataset import LayoutDataset, collate_layout +from enpu_train.data.synthetic import generate_synthetic_set +from enpu_train.engine.trainer import TrainConfig, build_model_from_cfg, train_loop +from enpu_train.export.weights import export_onnx, export_state_dict + + +def load_yaml(path: Path) -> dict: + return yaml.safe_load(path.read_text(encoding="utf-8")) or {} + + +def main(argv: list[str] | None = None) -> int: + ap = argparse.ArgumentParser(description="EnPu layout train (#95)") + ap.add_argument( + "--config", + type=Path, + default=ROOT / "configs" / "mvp_l2_l3.yaml", + ) + ap.add_argument("--device", type=str, default=None) + ap.add_argument("--epochs", type=int, default=None) + ap.add_argument("--skip-export", action="store_true") + args = ap.parse_args(argv) + + cfg_raw = load_yaml(args.config) + data_c = cfg_raw.get("data") or {} + train_c = cfg_raw.get("train") or {} + export_c = cfg_raw.get("export") or {} + tasks = list(cfg_raw.get("tasks") or ["l2", "l3"]) + + # resolve roots relative to config / train dir + roots = [] + for r in data_c.get("roots") or []: + p = Path(r) + if not p.is_absolute(): + p = (args.config.parent / p).resolve() + if not p.exists(): + p = (ROOT / r).resolve() + if not p.exists(): + p = (ROOT.parent / Path(r).name).resolve() # repo samples/layout + if not p.exists() and "samples" in r.replace("\\", "/"): + p = (ROOT.parent / "samples" / "layout").resolve() + roots.append(p) + + # always try repo samples/layout + repo_layout = ROOT.parent / "samples" / "layout" + if repo_layout.is_dir() and repo_layout not in roots: + roots.insert(0, repo_layout) + + synth_count = int(data_c.get("synth_count") or 0) + if synth_count > 0: + synth_dir = Path(data_c.get("synth_dir") or "data_cache/synth") + if not synth_dir.is_absolute(): + synth_dir = ROOT / synth_dir + print(f"generating {synth_count} synthetic samples → {synth_dir}") + generate_synthetic_set(synth_dir, n=synth_count, seed=42) + roots.append(synth_dir) + + page_size = tuple(data_c.get("page_size") or [384, 512]) + row_size = tuple(data_c.get("row_size") or [64, 256]) + l2_heat = int(data_c.get("l2_heat_len") or 128) + l3_heat = int(data_c.get("l3_heat_len") or 128) + + device = args.device or train_c.get("device") or "cpu" + if device == "cuda" and not torch.cuda.is_available(): + print("cuda not available, using cpu") + device = "cpu" + + out_dir = Path(train_c.get("out_dir") or "runs/mvp_l2_l3") + if not out_dir.is_absolute(): + out_dir = ROOT / out_dir + + tcfg = TrainConfig( + tasks=tasks, + epochs=int(args.epochs or train_c.get("epochs") or 3), + batch_size=int(train_c.get("batch_size") or 2), + lr=float(train_c.get("lr") or 1e-3), + weight_decay=float(train_c.get("weight_decay") or 1e-4), + l2_loss_weight=float(train_c.get("l2_loss_weight") or 1.0), + l3_loss_weight=float(train_c.get("l3_loss_weight") or 1.5), + device=device, + page_h=int(page_size[0]), + page_w=int(page_size[1]), + row_h=int(row_size[0]), + row_w=int(row_size[1]), + l2_heat_len=l2_heat, + l3_heat_len=l3_heat, + num_workers=int(train_c.get("num_workers") or 0), + out_dir=str(out_dir), + log_every=int(train_c.get("log_every") or 1), + ) + + print("data roots:", [str(r) for r in roots]) + ds = LayoutDataset( + roots, + page_size=(tcfg.page_h, tcfg.page_w), + row_size=(tcfg.row_h, tcfg.row_w), + l2_heat_len=tcfg.l2_heat_len, + l3_heat_len=tcfg.l3_heat_len, + tasks=tuple(tasks), + augment=bool(data_c.get("augment", False)), + ) + print(f"samples: {len(ds)}") + + n = len(ds) + indices = list(range(n)) + random.Random(0).shuffle(indices) + val_ratio = float(data_c.get("val_ratio") or 0.25) + n_val = max(1, int(n * val_ratio)) if n > 1 else 0 + if n_val > 0 and n_val < n: + val_idx = indices[:n_val] + train_idx = indices[n_val:] + else: + train_idx = indices + val_idx = indices[:1] # toy: evaluate on one sample + + train_loader = DataLoader( + Subset(ds, train_idx), + batch_size=tcfg.batch_size, + shuffle=True, + num_workers=tcfg.num_workers, + collate_fn=collate_layout, + ) + val_loader = DataLoader( + Subset(ds, val_idx), + batch_size=1, + shuffle=False, + num_workers=0, + collate_fn=collate_layout, + ) + + model = build_model_from_cfg(tcfg) + result = train_loop(model, train_loader, val_loader, tcfg) + print("done:", result) + + if not args.skip_export: + ckpt = Path(result["best_path"]) + sd_out = Path(export_c.get("state_dict") or (out_dir / "export" / "layout_net.pt")) + if not sd_out.is_absolute(): + sd_out = ROOT / sd_out + export_state_dict(ckpt, sd_out) + print("exported state_dict:", sd_out) + try: + onnx_dir = Path(export_c.get("onnx_dir") or (out_dir / "export" / "onnx")) + if not onnx_dir.is_absolute(): + onnx_dir = ROOT / onnx_dir + paths = export_onnx( + ckpt, + onnx_dir, + page_h=tcfg.page_h, + page_w=tcfg.page_w, + row_h=tcfg.row_h, + row_w=tcfg.row_w, + ) + print("exported onnx:", paths) + except Exception as e: + print("onnx export skipped:", e) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/train/scripts/viz_sample.py b/train/scripts/viz_sample.py new file mode 100644 index 0000000..d1e8453 --- /dev/null +++ b/train/scripts/viz_sample.py @@ -0,0 +1,32 @@ +#!/usr/bin/env python3 +"""Draw L1/L2 boxes + L3 splits on a layout sample (#95).""" + +from __future__ import annotations + +import argparse +import json +import sys +from pathlib import Path + +ROOT = Path(__file__).resolve().parents[1] +if str(ROOT) not in sys.path: + sys.path.insert(0, str(ROOT)) + +from enpu_train.viz import draw_layout_overlay + + +def main() -> int: + ap = argparse.ArgumentParser() + ap.add_argument("sample_dir", type=Path, help="dir with layout.json + image") + ap.add_argument("-o", "--out", type=Path, default=None) + args = ap.parse_args() + layout = json.loads((args.sample_dir / "layout.json").read_text(encoding="utf-8")) + img_name = (layout.get("image") or {}).get("path") or "image.png" + out = args.out or (args.sample_dir / "overlay_preview.png") + draw_layout_overlay(args.sample_dir / img_name, layout, out_path=out) + print("wrote", out) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/train/tests/test_train_smoke.py b/train/tests/test_train_smoke.py new file mode 100644 index 0000000..707808e --- /dev/null +++ b/train/tests/test_train_smoke.py @@ -0,0 +1,102 @@ +"""Smoke tests for train framework (#95) — no long GPU runs.""" + +from __future__ import annotations + +import json +import sys +from pathlib import Path + +import torch +from torch.utils.data import DataLoader + +ROOT = Path(__file__).resolve().parents[1] +if str(ROOT) not in sys.path: + sys.path.insert(0, str(ROOT)) + +from enpu_train.data.dataset import LayoutDataset, collate_layout, decode_peaks +from enpu_train.data.synthetic import make_synthetic_layout_sample +from enpu_train.engine.trainer import TrainConfig, build_model_from_cfg, train_loop +from enpu_train.metrics.layout_metrics import split_x_metrics +from enpu_train.viz import draw_layout_overlay + + +def test_synthetic_and_dataset(tmp_path: Path) -> None: + d = tmp_path / "S001" + layout = make_synthetic_layout_sample(d, sample_id="S001", seed=1) + assert (d / "layout.json").is_file() + assert len(layout["l2"]["systems"]) >= 1 + + ds = LayoutDataset( + tmp_path, + page_size=(192, 256), + row_size=(32, 128), + l2_heat_len=64, + l3_heat_len=64, + augment=False, + ) + assert len(ds) == 1 + item = ds[0] + assert item["page"].shape[0] == 3 + assert item["l2_heat"].shape[0] == 64 + assert len(item["row_images"]) >= 1 + + batch = collate_layout([item, item]) + assert batch["page"].shape[0] == 2 + assert batch["row_images"].shape[0] >= 1 + + +def test_split_metrics_and_peaks() -> None: + heat = torch.zeros(64) + heat[10] = 1.0 + heat[40] = 0.9 + peaks = decode_peaks(heat.numpy(), min_prominence=0.5, min_gap=5) + assert 10 in peaks and 40 in peaks + m = split_x_metrics([100.0, 200.0], [102.0, 198.0], max_dist=12) + assert m["split_count_exact"] == 1.0 + assert m["split_mean_abs_x_error"] <= 3.0 + + +def test_one_epoch_train(tmp_path: Path) -> None: + for i in range(3): + make_synthetic_layout_sample( + tmp_path / f"S{i:03d}", + sample_id=f"S{i:03d}", + seed=10 + i, + width=320, + height=400, + ) + ds = LayoutDataset( + tmp_path, + page_size=(192, 256), + row_size=(32, 128), + l2_heat_len=64, + l3_heat_len=64, + ) + loader = DataLoader(ds, batch_size=1, collate_fn=collate_layout, shuffle=True) + cfg = TrainConfig( + tasks=["l2", "l3"], + epochs=1, + batch_size=1, + device="cpu", + page_h=192, + page_w=256, + row_h=32, + row_w=128, + l2_heat_len=64, + l3_heat_len=64, + out_dir=str(tmp_path / "run"), + log_every=10, + ) + model = build_model_from_cfg(cfg) + result = train_loop(model, loader, loader, cfg) + assert Path(result["best_path"]).is_file() + hist = json.loads((tmp_path / "run" / "history.json").read_text(encoding="utf-8")) + assert hist[0]["train_loss"] == hist[0]["train_loss"] # not nan + + +def test_viz(tmp_path: Path) -> None: + d = tmp_path / "S" + layout = make_synthetic_layout_sample(d, sample_id="S", seed=0) + out = tmp_path / "ov.png" + draw_layout_overlay(d / "image.png", layout, out_path=out) + assert out.is_file()