"""승인 P-track paired source로 0.6 raster decoder만 보존형 미세조정한다.""" from __future__ import annotations import argparse import json import random import sys from pathlib import Path import numpy as np import torch from torch.utils.data import DataLoader, WeightedRandomSampler PROJECT_ROOT = Path(__file__).parents[1] SOURCE_ROOT = PROJECT_ROOT / "src" if str(SOURCE_ROOT) not in sys.path: sys.path.insert(0, str(SOURCE_ROOT)) if str(PROJECT_ROOT / "scripts") not in sys.path: sys.path.insert(0, str(PROJECT_ROOT / "scripts")) from math_grid_drawer.research.ink06_federation import ( FederatedPairedInk06Dataset, federation_provenance06, interpolate_state_dict06, load_product_federation06, resolve_training_device06, source_label_balanced_sampler06, ) from math_grid_drawer.research.math_ink_06 import MathInk06Engine, MathInk06Model from math_grid_drawer.research.math_ink_06 import virtual_raster_similarity06 from train_math_ink_06_candidate import _losses from train_math_ink_06_federated_online import _evaluate_sources, _partition, _source_subset def _macro_raster_score(metrics: dict[str, dict[str, float]]) -> float: """필요 변수: source별 raster 지표. 작동 원리: top-1을 우선하고 top-5를 보조하는 source-macro 선택 점수를 만든다.""" return float(np.mean([row["raster_top1"] + 0.25 * row["raster_top5"] for row in metrics.values()])) def _evaluate_geometry_sources( engine: MathInk06Engine, groups: dict[str, list[dict]], exact_to_index: dict[str, int], family_to_index: dict[str, int], batch_size: int, ) -> dict[str, dict[str, float]]: """필요 변수: source holdout·vectorizer. 작동 원리: label과 무관한 score-top1·top-4 raster 재구성도를 계산한다.""" reports = {} engine.model.eval() for source_id, records in groups.items(): loader = DataLoader( FederatedPairedInk06Dataset(records, exact_to_index, family_to_index), batch_size=batch_size, shuffle=False, num_workers=0, ) selected_total = oracle_total = samples = 0 with torch.inference_mode(): for _online, raster, _coordinates, _states, _target, _family, _source in loader: raster = raster.to(engine.device) output = engine.model.forward_raster(raster) similarity = virtual_raster_similarity06( output["coordinates"], raster, state_logits=output["state_logits"], size=32, sigma=0.025, ) selected = output["hypothesis_scores"].argmax(dim=1) batch_index = torch.arange(len(raster), device=engine.device) selected_total += float(similarity[batch_index, selected].sum()) oracle_total += float(similarity.amax(dim=1).sum()) samples += len(raster) reports[source_id] = { "samples": samples, "score_top1_similarity": selected_total / max(samples, 1), "top4_oracle_similarity": oracle_total / max(samples, 1), } return reports def _macro_geometry_score(metrics: dict[str, dict[str, float]]) -> float: """필요 변수: source별 재구성도. 작동 원리: 표본 수 편향 없이 top-4 기하 상한을 평균한다.""" return float(np.mean([row["top4_oracle_similarity"] for row in metrics.values()])) def _hard_label_sampler(records: list[dict], *, seed: int, hard_labels: set[str], multiplier: float) -> WeightedRandomSampler: """필요 변수: source-balanced record·hard label. 작동 원리: 기존 source/label 균형 위에서 공통 실패 기호만 제한적으로 재표집한다.""" base = source_label_balanced_sampler06(records, seed=seed, samples=len(records)) weights = base.weights.detach().clone() if multiplier < 1.0: raise ValueError("hard label multiplier는 1 이상이어야 합니다.") for index, record in enumerate(records): if str(record["label"]) in hard_labels: weights[index] *= multiplier generator = torch.Generator().manual_seed(seed) return WeightedRandomSampler(weights, len(records), replacement=True, generator=generator) def _split_sources(sources, *, seed: int, train_max: int, validation_max: int, test_max: int): """필요 변수: 승인 source·subset 상한. 작동 원리: 기존 federation과 동일한 writer/origin 분리로 train·validation·test를 고정한다.""" training: list[dict] = [] validation: dict[str, list[dict]] = {} test: dict[str, list[dict]] = {} for source_index, source in enumerate(sources): eligible = [row for row in source.records if row.get("eligible_for_training")] explicit_validation = [row for row in source.records if str(row.get("split")) in {"validation", "valid", "val"}] train_candidates = eligible if explicit_validation else [row for row in eligible if _partition(row) >= 2] validation_candidates = explicit_validation or [row for row in eligible if _partition(row) == 0] training.extend(_source_subset(train_candidates, train_max, seed + source_index)) validation[source.source_id] = _source_subset(validation_candidates, validation_max, seed + 20 + source_index) test_candidates = [row for row in source.records if row.get("split") == "test"] test[source.source_id] = _source_subset(test_candidates, test_max, seed + 40 + source_index) if any(not rows for rows in validation.values()) or any(not rows for rows in test.values()): raise ValueError("source validation/test partition이 비었습니다.") return training, validation, test def main() -> None: """필요 변수: 0.6 checkpoint·승인 federation. 작동 원리: online head를 고정하고 hard-label decoder 후보를 holdout으로 선택한다.""" parser = argparse.ArgumentParser(description="Train Math Ink 0.6 federated raster decoder") parser.add_argument("--checkpoint", type=Path, required=True) parser.add_argument("--registry", type=Path, default=PROJECT_ROOT / "research/dataset_registry.json") parser.add_argument("--source-registry", type=Path, default=PROJECT_ROOT / "research/math_ink_06_source_registry.json") parser.add_argument("--commercial", type=Path, default=PROJECT_ROOT / "research/data/external_trajectory_v1/commercial_ccby4.jsonl.gz") parser.add_argument("--hwrt", type=Path, default=PROJECT_ROOT / "research/data/open_pretrain/hwrt_expanded_v2/hwrt_expanded.jsonl.gz") parser.add_argument("--approval", type=Path, default=PROJECT_ROOT / "research/approvals/HWRT-ODBL-USE-APPROVAL-v1.json") parser.add_argument("--output", type=Path, required=True) parser.add_argument("--hard-labels", default="2,%,A,\\Delta,\\Leftrightarrow,\\mathbb{H},\\mu,\\varpi,p") parser.add_argument("--hard-label-multiplier", type=float, default=2.0) parser.add_argument("--max-train-per-source", type=int, default=1000) parser.add_argument("--max-validation-per-source", type=int, default=300) parser.add_argument("--max-test-per-source", type=int, default=500) parser.add_argument("--batch-size", type=int, default=32) parser.add_argument("--epochs", type=int, default=1) parser.add_argument("--learning-rate", type=float, default=1e-5) parser.add_argument("--cycle-weight", type=float, default=0.05) parser.add_argument("--preservation-weight", type=float, default=1.0) parser.add_argument("--reconstruction-weight", type=float, default=0.0) parser.add_argument("--reconstruction-size", type=int, default=32) parser.add_argument("--selection-mode", choices=("raster_classification", "geometry"), default="raster_classification") parser.add_argument( "--raster-architecture", choices=("fine_cross_attention_16x16_v6", "gated_fine_cross_attention_16x16_v7"), ) parser.add_argument("--classification-tolerance", type=float, default=0.005) parser.add_argument("--holdout-tolerance", type=float, default=0.005) parser.add_argument("--interpolation-alphas", default="0.125,0.25,0.5,1.0") parser.add_argument("--seed", type=int, default=17) parser.add_argument("--device", default="auto", help="auto|cpu|cuda[:index]") args = parser.parse_args() if min( args.cycle_weight, args.preservation_weight, args.reconstruction_weight, args.holdout_tolerance, args.classification_tolerance, ) < 0: raise ValueError("loss weight와 holdout tolerance는 0 이상이어야 합니다.") random.seed(args.seed) np.random.seed(args.seed) torch.manual_seed(args.seed) student_checkpoint = args.checkpoint if args.raster_architecture: # 기존 online/head를 그대로 두고 새 fine raster branch의 추가 weight만 초기화한다. source_payload = torch.load(args.checkpoint, map_location="cpu", weights_only=False) initialized_model = MathInk06Model( exact_classes=len(source_payload["exact_labels"]), family_classes=len(source_payload["family_labels"]), hidden_size=int(source_payload.get("hidden_size", 128)), hypotheses=int(source_payload.get("hypotheses", 4)), raster_architecture=args.raster_architecture, virtual_contract=str(source_payload.get("virtual_contract") or "legacy_v1"), ) initialized_model.load_state_dict(source_payload["state_dict"], strict=False) source_payload["state_dict"] = initialized_model.state_dict() source_payload["raster_architecture"] = args.raster_architecture args.output.mkdir(parents=True, exist_ok=True) student_checkpoint = args.output / "initialized_fine_checkpoint.pt" torch.save(source_payload, student_checkpoint) device = resolve_training_device06(args.device) student = MathInk06Engine(student_checkpoint, device=device) teacher = MathInk06Engine(args.checkpoint, device=str(student.device)) print(json.dumps({"device": str(student.device), "cuda": torch.cuda.is_available()}), flush=True) for parameter in student.model.parameters(): parameter.requires_grad_(False) trainable = [*student.model.raster_encoder.parameters(), *student.model.virtual_decoder.parameters()] if student.model.auxiliary_virtual_decoder is not None: trainable.extend(student.model.auxiliary_virtual_decoder.parameters()) for parameter in trainable: parameter.requires_grad_(True) exact_to_index = {label: index for index, label in enumerate(student.labels)} family_to_index = {label: index for index, label in enumerate(student.family_labels)} sources = load_product_federation06( registry_path=args.registry, commercial_path=args.commercial, hwrt_path=args.hwrt, approval_path=args.approval, allowed_labels=student.labels, source_registry_path=args.source_registry, ) training, validation_groups, test_groups = _split_sources( sources, seed=args.seed, train_max=args.max_train_per_source, validation_max=args.max_validation_per_source, test_max=args.max_test_per_source, ) hard_labels = {value.strip() for value in args.hard_labels.split(",") if value.strip()} unknown = hard_labels.difference(exact_to_index) if unknown: raise ValueError(f"378 vocabulary에 없는 hard label입니다: {sorted(unknown)}") sampler = _hard_label_sampler( training, seed=args.seed, hard_labels=hard_labels, multiplier=args.hard_label_multiplier, ) loader = DataLoader( FederatedPairedInk06Dataset(training, exact_to_index, family_to_index), batch_size=args.batch_size, sampler=sampler, num_workers=0, ) optimizer = torch.optim.AdamW(trainable, lr=args.learning_rate, weight_decay=1e-3) baseline_validation = _evaluate_sources( student, validation_groups, exact_to_index, family_to_index, args.batch_size, ) baseline_geometry = _evaluate_geometry_sources( student, validation_groups, exact_to_index, family_to_index, args.batch_size, ) baseline_test = _evaluate_sources(student, test_groups, exact_to_index, family_to_index, args.batch_size) best_score = ( _macro_geometry_score(baseline_geometry) if args.selection_mode == "geometry" else _macro_raster_score(baseline_validation) ) best_state = {key: value.detach().cpu().clone() for key, value in student.model.state_dict().items()} anchor_state = {key: value.clone() for key, value in best_state.items()} best_metrics = baseline_validation best_geometry = baseline_geometry best_epoch = 0 best_alpha = 0.0 alphas = tuple(float(value) for value in args.interpolation_alphas.split(",") if value.strip()) if not alphas or any(not 0 < value <= 1 for value in alphas): raise ValueError("interpolation alpha는 0보다 크고 1 이하여야 합니다.") history: list[dict] = [] for epoch in range(1, args.epochs + 1): student.model.train() totals: dict[str, float] = {} seen = 0 for online, raster, coordinates, states, target, family, _source in loader: online, raster, coordinates, states, target, family = [ value.to(student.device) for value in (online, raster, coordinates, states, target, family) ] optimizer.zero_grad(set_to_none=True) with torch.inference_mode(): teacher_output = teacher.model.forward_raster(raster) loss, components = _losses( student.model, online, raster, coordinates, states, target, family, cycle_weight=args.cycle_weight, online_weight=0.0, multi_target=True, teacher_output=teacher_output, preservation_weight=args.preservation_weight, reconstruction_weight=args.reconstruction_weight, reconstruction_size=args.reconstruction_size, ) loss.backward() torch.nn.utils.clip_grad_norm_(trainable, 1.0) optimizer.step() batch = len(target) seen += batch totals["loss"] = totals.get("loss", 0.0) + float(loss.detach()) * batch for key, value in components.items(): totals[key] = totals.get(key, 0.0) + value * batch trained_state = {key: value.detach().cpu().clone() for key, value in student.model.state_dict().items()} interpolation = [] for alpha in alphas: mixed = interpolate_state_dict06(anchor_state, trained_state, alpha=alpha) student.model.load_state_dict(mixed) metrics = _evaluate_sources(student, validation_groups, exact_to_index, family_to_index, args.batch_size) geometry = _evaluate_geometry_sources( student, validation_groups, exact_to_index, family_to_index, args.batch_size, ) score = ( _macro_geometry_score(geometry) if args.selection_mode == "geometry" else _macro_raster_score(metrics) ) tolerance = ( args.classification_tolerance if args.selection_mode == "geometry" else args.holdout_tolerance ) guard = all( metrics[source][metric] >= baseline_validation[source][metric] - tolerance for source in metrics for metric in ("raster_top1", "raster_top5") ) interpolation.append({ "alpha": alpha, "score": score, "holdout_guard": guard, "validation": metrics, "geometry": geometry, }) if guard and score > best_score: best_score, best_metrics, best_geometry, best_epoch, best_alpha = ( score, metrics, geometry, epoch, alpha ) best_state = {key: value.clone() for key, value in mixed.items()} student.model.load_state_dict(trained_state) row = { "epoch": epoch, "components": {key: value / max(seen, 1) for key, value in totals.items()}, "interpolation": interpolation, } history.append(row) print(json.dumps(row, ensure_ascii=False), flush=True) student.model.load_state_dict(best_state) final_test = _evaluate_sources(student, test_groups, exact_to_index, family_to_index, args.batch_size) final_test_geometry = _evaluate_geometry_sources( student, test_groups, exact_to_index, family_to_index, args.batch_size, ) payload = torch.load(student_checkpoint, map_location="cpu", weights_only=False) payload["state_dict"] = best_state payload["model_version"] = "aiflow-math-ink-0.6-federated-decoder1" provenance = federation_provenance06(sources, args.source_registry) payload.update(provenance) payload["federated_decoder"] = { "seed": args.seed, "hard_labels": sorted(hard_labels), "hard_label_multiplier": args.hard_label_multiplier, "selected_epoch": best_epoch, "selected_interpolation_alpha": best_alpha, "cycle_weight": args.cycle_weight, "preservation_weight": args.preservation_weight, "reconstruction_weight": args.reconstruction_weight, "reconstruction_size": args.reconstruction_size, "selection_mode": args.selection_mode, } payload["product_validation"] = False args.output.mkdir(parents=True, exist_ok=True) checkpoint = args.output / "math_ink_06_candidate.pt" torch.save(payload, checkpoint) report = { "checkpoint": checkpoint.name, "bytes": checkpoint.stat().st_size, "seed": args.seed, "device": str(student.device), **provenance, "train_samples": len(training), "source_count": len(sources), "hard_labels": sorted(hard_labels), "baseline_validation": baseline_validation, "baseline_geometry": baseline_geometry, "selected_validation": best_metrics, "selected_geometry": best_geometry, "selected_epoch": best_epoch, "selected_interpolation_alpha": best_alpha, "baseline_test": baseline_test, "test": final_test, "test_geometry": final_test_geometry, "test_delta": { source: {metric: final_test[source][metric] - baseline_test[source][metric] for metric in ( "online_top1", "online_top5", "raster_top1", "raster_top5", )} for source in final_test }, "history": history, "online_frozen": True, "product_validation": False, } (args.output / "federated_decoder_report.json").write_text( json.dumps(report, ensure_ascii=False, indent=2) + "\n", encoding="utf-8", ) print(json.dumps({key: report[key] for key in report if key not in {"history"}}, ensure_ascii=False, indent=2)) if __name__ == "__main__": main()