File size: 21,388 Bytes
bb698e6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
"""통과한 P Formula 3-seed teacher를 하나의 모바일 formula-adapter student로 증류한다."""

from __future__ import annotations

import argparse
from collections import Counter
from copy import deepcopy
from datetime import datetime, timezone
import json
from pathlib import Path
import random
import sys
from typing import Any, Sequence

import numpy as np
import torch
from torch import Tensor
from torch.utils.data import DataLoader, TensorDataset, WeightedRandomSampler

PROJECT_ROOT = Path(__file__).parents[1]
SOURCE_ROOT = PROJECT_ROOT / "src"
for path in (PROJECT_ROOT, SOURCE_ROOT):
    if str(path) not in sys.path:
        sys.path.insert(0, str(path))

from math_grid_drawer.research.external_corpus import read_jsonl
from math_grid_drawer.research.p_formula_dataset06 import (
    PFormulaTensorBatch06,
    materialize_p_formula_split06,
    p_formula_release_metrics06,
    p_formula_seed_gate06,
)
from math_grid_drawer.research.p_formula_gate06 import audit_p_formula_records06
from math_grid_drawer.research.skeleton_adapter06 import SkeletonTrajectoryAdapter06
from scripts.audit_math_ink_06_case_context import _load_model06
from scripts.train_math_ink_06_formula_adapter import _forward06, _targets06
from scripts.train_math_ink_06_p_formula_adapter import _file_sha25606


REQUIRED_SEEDS06 = frozenset({17, 31, 47})


def validate_p_formula_distillation_inputs06(
    summary: dict[str, Any],
    reports: Sequence[dict[str, Any]],
    *,
    data_sha256: str,
) -> list[dict[str, Any]]:
    """필요 변수: 3-seed summary·teacher report·현재 data hash. 작동 원리: 동일 P corpus와 전 seed 통과를 AND로 검증한다."""

    if summary.get("track") != "P_approved_formula_only":
        raise ValueError("P Formula distillation에는 P-track summary만 허용합니다.")
    decision = summary.get("decision") or {}
    if decision.get("student_distillation_allowed") is not True:
        raise ValueError("3-seed summary가 student distillation을 허용하지 않았습니다.")
    if str(summary.get("data_sha256") or "") != data_sha256:
        raise ValueError("현재 P Formula data SHA-256이 3-seed summary와 다릅니다.")
    if len(reports) != 3 or {int(report["seed"]) for report in reports} != REQUIRED_SEEDS06:
        raise ValueError("Teacher report는 seed 17·31·47이 정확히 하나씩 필요합니다.")
    ordered = sorted(reports, key=lambda report: int(report["seed"]))
    for report in ordered:
        if report.get("track") != "P_approved_formula_only":
            raise ValueError("R-track teacher를 P student에 증류할 수 없습니다.")
        if str(report.get("data_sha256") or "") != data_sha256:
            raise ValueError("Teacher report의 P Formula data SHA-256이 다릅니다.")
        if report.get("seed_gate", {}).get("passed") is not True:
            raise ValueError(f"seed {report['seed']} teacher가 개별 release gate를 통과하지 않았습니다.")
    return ordered


def ensemble_teacher_probability06(
    logits: Sequence[Tensor],
    *,
    temperature: float,
) -> Tensor:
    """필요 변수: seed별 동일 shape logits·temperature. 작동 원리: logit 평균 대신 확률 평균으로 teacher target을 만든다."""

    if not logits or temperature <= 0.0:
        raise ValueError("Teacher logit과 양수 temperature가 필요합니다.")
    shape = logits[0].shape
    if any(value.shape != shape for value in logits):
        raise ValueError("Teacher logit shape가 서로 다릅니다.")
    return torch.stack([
        (value / temperature).softmax(dim=1)
        for value in logits
    ]).mean(dim=0)


def _parse_args() -> argparse.Namespace:
    """필요 변수: P corpus·teacher reports/summary·student main adapter. 작동 원리: fail-closed distillation CLI를 만든다."""

    parser = argparse.ArgumentParser(description="Distill Math Ink 0.6 P formula student")
    parser.add_argument("--data", type=Path, required=True)
    parser.add_argument("--teacher-report", type=Path, action="append", required=True)
    parser.add_argument("--summary", type=Path, required=True)
    parser.add_argument("--student-adapter", type=Path, required=True)
    parser.add_argument("--output", type=Path, required=True)
    parser.add_argument("--seed", type=int, default=17)
    parser.add_argument("--epochs", type=int, default=20)
    parser.add_argument("--batch-size", type=int, default=128)
    parser.add_argument("--learning-rate", type=float, default=4e-4)
    parser.add_argument("--weight-decay", type=float, default=2e-3)
    parser.add_argument("--hidden-size", type=int, default=64)
    parser.add_argument("--temperature", type=float, default=2.0)
    parser.add_argument("--teacher-exact-weight", type=float, default=0.70)
    parser.add_argument("--teacher-family-weight", type=float, default=1.00)
    parser.add_argument("--hard-exact-weight", type=float, default=0.20)
    parser.add_argument("--hard-family-weight", type=float, default=0.50)
    parser.add_argument("--patience", type=int, default=5)
    parser.add_argument("--minimum-independent-sources", type=int, default=2)
    parser.add_argument("--distillation-regression-maximum-pp", type=float, default=1.0)
    parser.add_argument("--device", choices=("cuda", "cpu"), default="cuda")
    return parser.parse_args()


def _seed06(seed: int) -> None:
    """필요 변수: student seed. 작동 원리: Python·NumPy·PyTorch 초기화를 고정한다."""

    random.seed(seed)
    np.random.seed(seed)
    torch.manual_seed(seed)
    if torch.cuda.is_available():
        torch.cuda.manual_seed_all(seed)


def _resolve_project_path06(value: str | Path, *, parent: Path | None = None) -> Path:
    """필요 변수: checkpoint lineage 경로·선택 report parent. 작동 원리: 절대/상대 경로를 존재하는 실제 파일로 해석한다."""

    path = Path(value)
    candidates = [path] if path.is_absolute() else [
        *((parent / path,) if parent is not None else ()),
        PROJECT_ROOT / path,
    ]
    for candidate in candidates:
        if candidate.is_file():
            return candidate
    raise FileNotFoundError(f"checkpoint 경로를 찾을 수 없습니다: {value}")


def _load_teacher06(
    report: dict[str, Any],
    report_path: Path,
    *,
    device: torch.device,
) -> tuple[Any, torch.nn.Module, SkeletonTrajectoryAdapter06, tuple[str, ...]]:
    """필요 변수: 통과 teacher report/path. 작동 원리: base→shared online→P formula adapter 순서로 합성한다."""

    formula_checkpoint = _resolve_project_path06(
        str(report["checkpoint"]),
        parent=report_path.parent,
    )
    payload = torch.load(formula_checkpoint, map_location="cpu", weights_only=False)
    if payload.get("track") != "P_approved_formula_only":
        raise ValueError("P Formula teacher checkpoint track이 올바르지 않습니다.")
    if payload.get("seed_gate_passed") is not True:
        raise ValueError("개별 gate를 통과하지 않은 teacher checkpoint입니다.")
    if str(payload.get("data_sha256") or "") != str(report["data_sha256"]):
        raise ValueError("Teacher checkpoint/report data SHA-256이 다릅니다.")
    online_path = _resolve_project_path06(str(payload["online_adapter"]))
    base_path = _resolve_project_path06(str(payload["base_checkpoint"]))
    engine, online_adapter = _load_model06(base_path, online_path, device)
    formula_adapter = SkeletonTrajectoryAdapter06(
        hidden_size=int(payload["hidden_size"]),
    ).to(device)
    formula_adapter.load_state_dict(payload["state_dict"])
    engine.model.eval()
    online_adapter.eval()
    formula_adapter.eval()
    return engine, online_adapter, formula_adapter, tuple(str(label) for label in engine.labels)


def _teacher_targets06(
    teachers: Sequence[tuple[Any, torch.nn.Module, SkeletonTrajectoryAdapter06]],
    features: Tensor,
    *,
    temperature: float,
    device: torch.device,
    batch_size: int,
) -> tuple[Tensor, Tensor]:
    """필요 변수: 세 teacher·한 split feature. 작동 원리: seed별 exact/family probability를 CPU에서 평균한다."""

    exact_rows, family_rows = [], []
    for engine, online_adapter, formula_adapter in teachers:
        exact, family = _forward06(
            engine.model,
            online_adapter,
            formula_adapter,
            features,
            device=device,
            batch_size=batch_size,
        )
        exact_rows.append(exact)
        family_rows.append(family)
    return (
        ensemble_teacher_probability06(exact_rows, temperature=temperature),
        ensemble_teacher_probability06(family_rows, temperature=temperature),
    )


def _balanced_loader06(
    batch: PFormulaTensorBatch06,
    tensors: Sequence[Tensor],
    *,
    batch_size: int,
    seed: int,
) -> DataLoader:
    """필요 변수: P batch·학습 tensor. 작동 원리: source×label 역제곱근 sampler로 distillation batch를 만든다."""

    exact_targets = tensors[0]
    label_counts = Counter(int(value) for value in exact_targets.tolist())
    source_counts = Counter(batch.source_ids)
    weights = torch.tensor([
        1.0 / (
            max(label_counts[int(label)], 1) ** 0.5
            * max(source_counts[source], 1) ** 0.5
        )
        for label, source in zip(
            exact_targets.tolist(),
            batch.source_ids,
            strict=True,
        )
    ])
    sampler = WeightedRandomSampler(
        weights,
        num_samples=len(weights),
        replacement=True,
        generator=torch.Generator().manual_seed(seed),
    )
    return DataLoader(
        TensorDataset(batch.features, *tensors),
        batch_size=batch_size,
        sampler=sampler,
    )


def _release_metrics06(
    logits: Tensor,
    targets: Tensor,
    batch: PFormulaTensorBatch06,
    labels: Sequence[str],
) -> dict[str, Any]:
    """필요 변수: student/teacher exact logit·split metadata. 작동 원리: 공통 P release metric을 호출한다."""

    return p_formula_release_metrics06(
        logits,
        targets,
        labels=labels,
        writer_ids=batch.writer_ids,
        source_ids=batch.source_ids,
        timestamp_missing=batch.timestamp_missing,
        pressure_missing=batch.pressure_missing,
    )


def main() -> None:
    """필요 변수: 통과한 세 teacher와 동일 P corpus. 작동 원리: 하나의 formula adapter student를 학습하고 test regression을 판정한다."""

    args = _parse_args()
    device = torch.device(args.device)
    if device.type == "cuda" and not torch.cuda.is_available():
        raise RuntimeError("CUDA distillation을 요청했지만 사용할 수 없습니다.")
    _seed06(args.seed)
    data_sha256 = _file_sha25606(args.data)
    summary = json.loads(args.summary.read_text(encoding="utf-8"))
    reports = [
        json.loads(path.read_text(encoding="utf-8"))
        for path in args.teacher_report
    ]
    ordered_reports = validate_p_formula_distillation_inputs06(
        summary,
        reports,
        data_sha256=data_sha256,
    )
    report_paths = {
        int(json.loads(path.read_text(encoding="utf-8"))["seed"]): path
        for path in args.teacher_report
    }
    loaded = [
        _load_teacher06(
            report,
            report_paths[int(report["seed"])],
            device=device,
        )
        for report in ordered_reports
    ]
    label_contracts = {labels for *_modules, labels in loaded}
    if len(label_contracts) != 1:
        raise ValueError("세 teacher의 exact vocabulary가 다릅니다.")
    labels = next(iter(label_contracts))
    teachers = [(engine, online, formula) for engine, online, formula, _labels in loaded]

    records = list(read_jsonl(args.data))
    audit = audit_p_formula_records06(
        records,
        minimum_independent_sources=args.minimum_independent_sources,
    )
    if not audit["eligible_for_product_evaluation"]:
        raise ValueError("현재 P Formula corpus가 product preflight를 통과하지 못했습니다.")
    split_records = {
        split: [record for record in records if str(record["split"]) == split]
        for split in ("training", "validation", "test")
    }
    batches = {
        split: materialize_p_formula_split06(values, allowed_labels=labels)
        for split, values in split_records.items()
    }
    targets = {
        split: _targets06(batch.truths, labels, loaded[0][0].family_labels)
        for split, batch in batches.items()
    }
    teacher_targets = {
        split: _teacher_targets06(
            teachers,
            batch.features,
            temperature=args.temperature,
            device=device,
            batch_size=args.batch_size,
        )
        for split, batch in batches.items()
    }

    student_adapter_path = _resolve_project_path06(args.student_adapter)
    student_payload = torch.load(
        student_adapter_path,
        map_location="cpu",
        weights_only=False,
    )
    student_base = _resolve_project_path06(str(student_payload["base_checkpoint"]))
    student_engine, student_online = _load_model06(
        student_base,
        student_adapter_path,
        device,
    )
    if tuple(str(label) for label in student_engine.labels) != labels:
        raise ValueError("Student main vocabulary가 teacher와 다릅니다.")
    for parameter in student_engine.model.parameters():
        parameter.requires_grad_(False)
    for parameter in student_online.parameters():
        parameter.requires_grad_(False)
    student_formula = SkeletonTrajectoryAdapter06(
        hidden_size=args.hidden_size,
    ).to(device)
    train_exact, train_family = targets["training"]
    train_teacher_exact, train_teacher_family = teacher_targets["training"]
    loader = _balanced_loader06(
        batches["training"],
        (
            train_exact,
            train_family,
            train_teacher_exact,
            train_teacher_family,
        ),
        batch_size=args.batch_size,
        seed=args.seed,
    )
    optimizer = torch.optim.AdamW(
        student_formula.parameters(),
        lr=args.learning_rate,
        weight_decay=args.weight_decay,
    )
    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(
        optimizer,
        T_max=max(args.epochs, 1),
        eta_min=args.learning_rate * 0.1,
    )
    best_key = (-1.0, -1.0)
    best_state: dict[str, Tensor] | None = None
    best_epoch = 0
    stale = 0
    history = []
    for epoch in range(1, args.epochs + 1):
        student_formula.train()
        losses = []
        for (
            features,
            exact_target,
            family_target,
            teacher_exact,
            teacher_family,
        ) in loader:
            features = features.to(device)
            exact_target = exact_target.to(device)
            family_target = family_target.to(device)
            teacher_exact = teacher_exact.to(device)
            teacher_family = teacher_family.to(device)
            optimizer.zero_grad(set_to_none=True)
            with torch.no_grad():
                online = student_online(features)
            exact_logits, family_logits = student_engine.model.classify_trajectory(
                student_formula(online),
            )
            temperature = args.temperature
            exact_distill = torch.nn.functional.kl_div(
                (exact_logits / temperature).log_softmax(dim=1),
                teacher_exact,
                reduction="batchmean",
            ) * temperature ** 2
            family_distill = torch.nn.functional.kl_div(
                (family_logits / temperature).log_softmax(dim=1),
                teacher_family,
                reduction="batchmean",
            ) * temperature ** 2
            loss = (
                args.teacher_exact_weight * exact_distill
                + args.teacher_family_weight * family_distill
                + args.hard_exact_weight
                * torch.nn.functional.cross_entropy(exact_logits, exact_target)
                + args.hard_family_weight
                * torch.nn.functional.cross_entropy(family_logits, family_target)
            )
            loss.backward()
            torch.nn.utils.clip_grad_norm_(student_formula.parameters(), 2.0)
            optimizer.step()
            losses.append(float(loss.detach()))
        scheduler.step()
        validation_logits = _forward06(
            student_engine.model,
            student_online,
            student_formula,
            batches["validation"].features,
            device=device,
            batch_size=args.batch_size,
        )
        validation_metrics = _release_metrics06(
            validation_logits[0],
            targets["validation"][0],
            batches["validation"],
            labels,
        )
        row = {
            "epoch": epoch,
            "loss": sum(losses) / max(len(losses), 1),
            "validation": validation_metrics,
        }
        history.append(row)
        print(json.dumps(row, ensure_ascii=False), flush=True)
        key = (
            float(validation_metrics["visual_family_top1"]),
            float(validation_metrics["exact_top1"]),
        )
        if key > best_key:
            best_key = key
            best_epoch = epoch
            best_state = deepcopy({
                name: value.detach().cpu()
                for name, value in student_formula.state_dict().items()
            })
            stale = 0
        else:
            stale += 1
        if stale >= args.patience:
            break
    if best_state is None:
        raise RuntimeError("Distilled student checkpoint가 선택되지 않았습니다.")
    student_formula.load_state_dict(best_state)
    student_test_logits = _forward06(
        student_engine.model,
        student_online,
        student_formula,
        batches["test"].features,
        device=device,
        batch_size=args.batch_size,
    )[0]
    student_test = _release_metrics06(
        student_test_logits,
        targets["test"][0],
        batches["test"],
        labels,
    )
    teacher_exact_probability = teacher_targets["test"][0]
    teacher_test = _release_metrics06(
        teacher_exact_probability.clamp_min(1e-9).log(),
        targets["test"][0],
        batches["test"],
        labels,
    )
    seed_gate = p_formula_seed_gate06(student_test)
    regression = {
        metric: (
            float(teacher_test[metric]) - float(student_test[metric])
        ) * 100.0
        for metric in ("exact_top1", "exact_top5", "visual_family_top1")
    }
    regression_passed = all(
        drop <= args.distillation_regression_maximum_pp
        for drop in regression.values()
    )
    distillation_gate_passed = bool(seed_gate["passed"] and regression_passed)
    args.output.mkdir(parents=True, exist_ok=True)
    checkpoint = args.output / "p_formula_student_adapter.pt"
    torch.save({
        "schema": "aiflow-math-ink-06-p-formula-student-v1",
        "state_dict": best_state,
        "hidden_size": args.hidden_size,
        "student_base_checkpoint": str(student_base),
        "student_online_adapter": str(student_adapter_path),
        "teacher_seeds": [17, 31, 47],
        "data_sha256": data_sha256,
        "selected_epoch": best_epoch,
        "distillation_gate_passed": distillation_gate_passed,
        "track": "P_approved_formula_only",
        "teacher_weights_embedded": False,
        "litert_exported": False,
        "product_validation": False,
    }, checkpoint)
    report = {
        "experiment": "P-MATH-INK-06-FORMULA-STUDENT-DISTILL-001",
        "generated_at": datetime.now(timezone.utc).isoformat(),
        "student_seed": args.seed,
        "teacher_seeds": [17, 31, 47],
        "data": str(args.data),
        "data_sha256": data_sha256,
        "preflight": audit,
        "temperature": args.temperature,
        "loss_weights": {
            "teacher_exact": args.teacher_exact_weight,
            "teacher_family": args.teacher_family_weight,
            "hard_exact": args.hard_exact_weight,
            "hard_family": args.hard_family_weight,
        },
        "selected_epoch": best_epoch,
        "teacher_test": teacher_test,
        "student_test": student_test,
        "student_seed_gate": seed_gate,
        "teacher_to_student_drop_pp": regression,
        "distillation_regression_maximum_pp": args.distillation_regression_maximum_pp,
        "distillation_regression_passed": regression_passed,
        "distillation_gate_passed": distillation_gate_passed,
        "history": history,
        "checkpoint": checkpoint.name,
        "checkpoint_bytes": checkpoint.stat().st_size,
        "teacher_weights_embedded": False,
        "track": "P_approved_formula_only",
        "litert_exported": False,
        "product_validation": False,
        "next_gate": "composite torch.export→LiteRT parity→Android 3-tier benchmark",
    }
    (args.output / "report.json").write_text(
        json.dumps(report, ensure_ascii=False, indent=2) + "\n",
        encoding="utf-8",
    )
    print(json.dumps({
        "student_test": student_test,
        "teacher_test": teacher_test,
        "distillation_gate_passed": distillation_gate_passed,
        "checkpoint": str(checkpoint),
        "product_validation": False,
    }, ensure_ascii=False, indent=2))


if __name__ == "__main__":
    main()