#!/usr/bin/env python3 """Fine-tune Brainmu from BaseNet reconstructions with explicit gray reconstruction loss. All pretrained Brainmu, ViT, VAE, connector, understanding-expert and norm parameters remain frozen. LoRA is injected only into nn.Linear modules whose logical name contains `_moe_gen` (generation-side attention and MLP experts). """ from __future__ import annotations import argparse import functools import json import math import os import random import shutil import sys import time from dataclasses import asdict, dataclass from pathlib import Path from typing import Dict, Iterable, List, Tuple import cv2 import numpy as np import torch import torch.nn as nn import torch.nn.functional as F from accelerate import init_empty_weights, load_checkpoint_and_dispatch from PIL import Image, ImageDraw from safetensors.torch import save_file from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import ( CheckpointImpl, apply_activation_checkpointing, checkpoint_wrapper, ) @dataclass class AdapterSpec: rank: int alpha: float dropout: float target_rule: str target_modules: List[str] trainable_parameters: int total_parameters: int class LoRALinear(nn.Module): def __init__(self, base: nn.Linear, rank: int, alpha: float, dropout: float): super().__init__() self.base = base self.base.requires_grad_(False) self.rank = rank self.alpha = alpha self.scaling = alpha / rank self.dropout = nn.Dropout(dropout) if dropout > 0 else nn.Identity() # Keep FP32 master weights; autocast executes the matmuls in BF16. self.lora_A = nn.Parameter(torch.empty(rank, base.in_features, device=base.weight.device, dtype=torch.float32)) self.lora_B = nn.Parameter(torch.zeros(base.out_features, rank, device=base.weight.device, dtype=torch.float32)) nn.init.kaiming_uniform_(self.lora_A, a=math.sqrt(5)) def forward(self, x: torch.Tensor) -> torch.Tensor: base_out = self.base(x) update = F.linear(F.linear(self.dropout(x), self.lora_A), self.lora_B) return base_out + update * self.scaling def set_seed(seed: int) -> None: random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) def resolve_parent(root: nn.Module, dotted: str) -> Tuple[nn.Module, str]: parts = dotted.split(".") parent = root for part in parts[:-1]: parent = parent[int(part)] if part.isdigit() else getattr(parent, part) return parent, parts[-1] def inject_generation_expert_lora( model: nn.Module, rank: int, alpha: float, dropout: float ) -> Dict[str, LoRALinear]: candidates = [ name for name, module in model.named_modules() if "_moe_gen" in name and isinstance(module, nn.Linear) ] if not candidates: raise RuntimeError("no generation-expert nn.Linear modules matched `_moe_gen`") adapters: Dict[str, LoRALinear] = {} for name in candidates: parent, leaf = resolve_parent(model, name) base = getattr(parent, leaf) wrapped = LoRALinear(base, rank=rank, alpha=alpha, dropout=dropout) setattr(parent, leaf, wrapped) adapters[name] = wrapped return adapters def adapter_parameters(adapters: Dict[str, LoRALinear]) -> List[nn.Parameter]: return [p for module in adapters.values() for p in (module.lora_A, module.lora_B)] def save_adapter(adapters: Dict[str, LoRALinear], path: Path) -> None: state = {} for name, module in adapters.items(): state[f"{name}.lora_A"] = module.lora_A.detach().cpu().contiguous() state[f"{name}.lora_B"] = module.lora_B.detach().cpu().contiguous() path.parent.mkdir(parents=True, exist_ok=True) save_file(state, str(path)) def ssim_gray(a: np.ndarray, b: np.ndarray) -> float: c1, c2 = 0.01**2, 0.03**2 mu_a = cv2.GaussianBlur(a, (11, 11), 1.5) mu_b = cv2.GaussianBlur(b, (11, 11), 1.5) sigma_a = cv2.GaussianBlur(a * a, (11, 11), 1.5) - mu_a * mu_a sigma_b = cv2.GaussianBlur(b * b, (11, 11), 1.5) - mu_b * mu_b sigma_ab = cv2.GaussianBlur(a * b, (11, 11), 1.5) - mu_a * mu_b score = ((2 * mu_a * mu_b + c1) * (2 * sigma_ab + c2)) / ( (mu_a * mu_a + mu_b * mu_b + c1) * (sigma_a + sigma_b + c2) ) return float(score.mean()) def image_metrics(pred: Image.Image, target: Image.Image) -> Dict[str, float]: p = np.asarray(pred.convert("L"), dtype=np.float32) / 255.0 t = np.asarray(target.convert("L").resize(pred.size), dtype=np.float32) / 255.0 mse = float(np.mean((p - t) ** 2)) return {"psnr_db": float(-10 * np.log10(max(mse, 1e-12))), "ssim": ssim_gray(p, t)} def main() -> None: ap = argparse.ArgumentParser() ap.add_argument("--brainmu-repo", type=Path, default=Path(__file__).resolve().parents[1]/"vendor/Brainmu") ap.add_argument("--model-path", type=Path, required=True) ap.add_argument("--train-dataset-root", type=Path, required=True) ap.add_argument("--test-dataset-root", type=Path, required=True) ap.add_argument("--parquet-dir", type=Path, required=True) ap.add_argument("--output-dir", type=Path, required=True) ap.add_argument("--steps", type=int, default=200) ap.add_argument("--lr", type=float, default=5e-4) ap.add_argument("--rank", type=int, default=8) ap.add_argument("--alpha", type=float, default=16.0) ap.add_argument("--dropout", type=float, default=0.0) ap.add_argument("--seed", type=int, default=20260829) ap.add_argument("--eval-every", type=int, default=200) ap.add_argument("--eval-steps", type=int, default=12) ap.add_argument("--train-eval-count", type=int, default=40) ap.add_argument("--test-eval-count", type=int, default=-1) ap.add_argument("--flow-weight", type=float, default=0.1) ap.add_argument("--gray-weight", type=float, default=1.0) args = ap.parse_args() for name in ["brainmu_repo","model_path","train_dataset_root","test_dataset_root","parquet_dir","output_dir"]: setattr(args,name,getattr(args,name).resolve()) args.output_dir.mkdir(parents=True, exist_ok=True) sys.path.insert(0, str(args.brainmu_repo.resolve())) os.chdir(args.brainmu_repo) from data.data_utils import add_special_tokens from data.dataset_base import DataConfig, PackedDataset, SimpleCustomBatch from data.dataset_info import DATASET_INFO from data.transforms import ImageTransform from inferencer import InterleaveInferencer from modeling.autoencoder import load_ae from modeling.brainmu import ( Brainmu, BrainmuConfig, Qwen2Config, Qwen2ForCausalLM, SiglipVisionConfig, SiglipVisionModel, ) from modeling.brainmu.qwen2_navit import Qwen2MoTDecoderLayer from modeling.qwen2 import Qwen2Tokenizer set_seed(args.seed) device = torch.device("cuda:0") torch.cuda.set_device(device) torch.backends.cuda.matmul.allow_tf32 = True model_path = args.model_path.resolve() llm_config = Qwen2Config.from_json_file(str(model_path / "llm_config.json")) llm_config.qk_norm = True llm_config.tie_word_embeddings = False llm_config.layer_module = "Qwen2MoTDecoderLayer" vit_config = SiglipVisionConfig.from_json_file(str(model_path / "vit_config.json")) vit_config.rope = False vit_config.num_hidden_layers -= 1 vae_model, vae_config = load_ae(local_path=str(model_path / "ae.safetensors")) config = BrainmuConfig( visual_gen=True, visual_und=True, llm_config=llm_config, vit_config=vit_config, vae_config=vae_config, vit_max_num_patch_per_side=70, connector_act="gelu_pytorch_tanh", latent_patch_size=2, max_latent_size=64, timestep_shift=1.0, ) with init_empty_weights(): language_model = Qwen2ForCausalLM(llm_config) vit_model = SiglipVisionModel(vit_config) model = Brainmu(language_model, vit_model, config) model.vit_model.vision_model.embeddings.convert_conv2d_to_linear(vit_config, meta=True) print("LOAD_MODEL_BEGIN", flush=True) model = load_checkpoint_and_dispatch( model, checkpoint=str(model_path / "ema.safetensors"), device_map={"": 0}, dtype=torch.bfloat16, # Brainmu's inference helpers create index tensors on CPU. The official # app enables hooks so those inputs follow the module device. force_hooks=True, ) model.requires_grad_(False) model.eval() # Keep the VAE in FP32 like Brainmu's official app. Its inference helper # supplies FP32 image tensors, while training pre-encoding is autocast. vae_model = vae_model.to(device=device, dtype=torch.float32).eval().requires_grad_(False) vae_param = next(vae_model.parameters()) vae_device, vae_dtype = vae_param.device, vae_param.dtype original_vae_encode = vae_model.encode original_vae_decode = vae_model.decode vae_model.encode = lambda x: original_vae_encode(x.to(device=vae_device, dtype=vae_dtype)) vae_model.decode = lambda z: original_vae_decode(z.to(device=vae_device, dtype=vae_dtype)) print(f"VAE_AUDIT device={vae_device} dtype={vae_dtype}", flush=True) print("LOAD_MODEL_DONE", flush=True) tokenizer = Qwen2Tokenizer.from_pretrained(str(model_path)) tokenizer, new_token_ids, _ = add_special_tokens(tokenizer) adapters = inject_generation_expert_lora(model, args.rank, args.alpha, args.dropout) trainable = adapter_parameters(adapters) trainable_ids = {id(p) for p in trainable} leaked = [name for name, p in model.named_parameters() if p.requires_grad and id(p) not in trainable_ids] frozen_lora = [name for name, p in model.named_parameters() if "lora_" in name and not p.requires_grad] if leaked or frozen_lora: raise RuntimeError(f"trainable audit failed: leaked={leaked}, frozen_lora={frozen_lora}") trainable_count = sum(p.numel() for p in trainable) total_count = sum(p.numel() for p in model.parameters()) spec = AdapterSpec( rank=args.rank, alpha=args.alpha, dropout=args.dropout, target_rule="nn.Linear and logical module name contains `_moe_gen`", target_modules=list(adapters), trainable_parameters=trainable_count, total_parameters=total_count, ) (args.output_dir / "adapter_config.json").write_text(json.dumps(asdict(spec), indent=2), encoding="utf-8") (args.output_dir / "trainable_audit.json").write_text( json.dumps( { "matched_modules": len(adapters), "trainable_parameters": trainable_count, "trainable_fraction": trainable_count / total_count, "only_generation_expert_lora": not leaked and not frozen_lora, "leaked_parameters": leaked, }, indent=2, ), encoding="utf-8", ) print( f"LORA_AUDIT matched={len(adapters)} trainable={trainable_count} total={total_count} " f"fraction={trainable_count/total_count:.8f}", flush=True, ) parquet_file = (args.parquet_dir / "reds_basenet_train2000.parquet").resolve() parquet_info = (args.parquet_dir / "parquet_info.json").resolve() DATASET_INFO["unified_edit"]["reds_basenet_train2000"] = { "data_dir": str(args.parquet_dir.resolve()), "num_files": 1, "num_total_samples": 2000, "parquet_info_path": str(parquet_info), } grouped = { "unified_edit": { "dataset_names": ["reds_basenet_train2000"], "image_transform_args": {"image_stride": 16, "max_image_size": 400, "min_image_size": 256}, "vit_image_transform_args": {"image_stride": 14, "max_image_size": 392, "min_image_size": 252}, "is_mandatory": False, "num_used_data": [1], "weight": 1.0, } } data_config = DataConfig(grouped_datasets=grouped) data_config.vae_image_downsample = 16 data_config.max_latent_size = 64 data_config.text_cond_dropout_prob = 0.0 data_config.vae_cond_dropout_prob = 0.0 data_config.vit_cond_dropout_prob = 0.0 data_config.vit_patch_size = 14 data_config.max_num_patch_per_side = 70 packed = PackedDataset( data_config, tokenizer=tokenizer, special_tokens=new_token_ids, local_rank=0, world_size=1, num_workers=1, expected_num_tokens=1100, max_num_tokens_per_sample=4096, max_num_tokens=4096, max_buffer_size=8, prefer_buffer_before=0, interpolate_pos=False, use_flex=False, data_status=None, ) packed.set_epoch(args.seed) raw_iter = iter(packed) batches = [SimpleCustomBatch([next(raw_iter)]).cuda(device) for _ in range(2000)] seen = [b.batch_data_indexes for b in batches] (args.output_dir / "cached_batch_indexes.json").write_text(json.dumps(seen, indent=2), encoding="utf-8") row_ids = [int(x[0]["data_indexes"][1]) for x in seen] if sorted(row_ids) != list(range(2000)): raise RuntimeError(f"train parquet coverage audit failed: {row_ids}") prepared_batches = [] with torch.no_grad(), torch.autocast("cuda", dtype=torch.bfloat16): for batch in batches: data = batch.to_dict() data.pop("batch_data_indexes", None) target_pixels = data.pop("padded_images") data["padded_latent"] = vae_model.encode(target_pixels) target_gray = ((target_pixels.float() + 1) * 0.5 * target_pixels.new_tensor([0.299, 0.587, 0.114])[None, :, None, None]).sum(1, keepdim=True).clamp(0, 1) prepared_batches.append((data, target_gray)) flow_capture = {} def capture_llm2vae(_module, _inputs, output): flow_capture["pred"] = output model.llm2vae.register_forward_hook(capture_llm2vae) def pack_clean_latent(data: dict) -> torch.Tensor: p = model.latent_patch_size; packed = [] for latent, (h, w) in zip(data["padded_latent"], data["patchified_vae_latent_shapes"]): x = latent[:, :h*p, :w*p].reshape(model.latent_channel, h, p, w, p) packed.append(torch.einsum("chpwq->hwpqc", x).reshape(-1, p*p*model.latent_channel)) return torch.cat(packed, 0) def unpack_generated_latent(tokens: torch.Tensor, shapes, timesteps) -> Tuple[List[torch.Tensor], List[int]]: p = model.latent_patch_size; out = []; image_ids = []; all_offset = 0; gen_offset = 0 for image_id, (h, w) in enumerate(shapes): n = h*w; block = timesteps[all_offset:all_offset+n] > 0; all_offset += n if bool(block.all()): x = tokens[gen_offset:gen_offset+n].reshape(h, w, p, p, model.latent_channel); gen_offset += n out.append(torch.einsum("hwpqc->chpwq", x).reshape(model.latent_channel, h*p, w*p)); image_ids.append(image_id) elif bool(block.any()): raise RuntimeError("mixed clean/noisy timesteps inside one image are unsupported") if gen_offset != len(tokens): raise RuntimeError(f"generated latent unpack mismatch {gen_offset}!={len(tokens)}") return out, image_ids def batch_loss(prepared, noise_seed: int, backward: bool = False): data, target_gray = prepared torch.manual_seed(noise_seed) rng_state = torch.cuda.get_rng_state(device) with torch.autocast("cuda", dtype=torch.bfloat16): loss_dict = model(**data) mse = loss_dict["mse"] flow_loss = mse.mean(dim=-1).sum() / max(1, len(data["mse_loss_indexes"])) pred_v = flow_capture.pop("pred") clean = pack_clean_latent(data) torch.cuda.set_rng_state(rng_state, device) noise = torch.randn_like(clean) t = torch.sigmoid(data["packed_timesteps"]) t = model.timestep_shift*t/(1+(model.timestep_shift-1)*t) has = t > 0 x_t = (1-t[:,None])*clean + t[:,None]*noise x0_hat = x_t[has] - t[has,None]*pred_v decoded, target_image_ids = unpack_generated_latent(x0_hat, data["patchified_vae_latent_shapes"], t) gray_losses = [] for z, image_id in zip(decoded, target_image_ids): tgt = target_gray[image_id] rgb = ((vae_model.decode(z[None]).float()+1)*0.5).clamp(0,1) lum = (rgb*rgb.new_tensor([0.299,0.587,0.114])[None,:,None,None]).sum(1,keepdim=True) if lum.shape[-2:] != tgt.shape[-2:]: tgt = F.interpolate(tgt[None], size=lum.shape[-2:], mode="bilinear", align_corners=False)[0] gray_losses.append(F.l1_loss(lum, tgt[None])) gray_loss = torch.stack(gray_losses).mean() total = args.flow_weight*flow_loss + args.gray_weight*gray_loss return total, flow_loss.detach(), gray_loss.detach() flow_probe_ids = np.linspace(0, len(prepared_batches) - 1, args.train_eval_count, dtype=int).tolist() flow_probe_batches = [prepared_batches[i] for i in flow_probe_ids] @torch.no_grad() def fixed_train_flow_mse() -> float: # Brainmu dispatches Qwen2 to its packed training forward only in train mode. # All base dropout is deterministic for this recipe and LoRA dropout is 0. model.train() vals = [float(batch_loss(d, args.seed + 10_000 + i)[0].item()) for i, d in enumerate(flow_probe_batches)] return float(np.mean(vals)) vae_transform = ImageTransform(400, 256, 16) vit_transform = ImageTransform(392, 252, 14) inferencer = InterleaveInferencer(model, vae_model, tokenizer, vae_transform, vit_transform, new_token_ids) train_records = [json.loads(x) for x in (args.train_dataset_root / "manifest.jsonl").read_text().splitlines()] test_records = [json.loads(x) for x in (args.test_dataset_root / "manifest.jsonl").read_text().splitlines()] if len(train_records) != 2000 or len(test_records) != 200: raise RuntimeError(f"expected train=2000/validation=200, got {len(train_records)}/{len(test_records)}") train_probe_ids = np.linspace(0, len(train_records) - 1, args.train_eval_count, dtype=int).tolist() eval_records = [] for idx in train_probe_ids: rec = dict(train_records[idx]); rec["split"] = "train"; rec["_root"] = str(args.train_dataset_root) eval_records.append(rec) selected_test_records = test_records if args.test_eval_count < 0 else test_records[:args.test_eval_count] for source in selected_test_records: rec = dict(source); rec["split"] = "test"; rec["_root"] = str(args.test_dataset_root) eval_records.append(rec) (args.output_dir / "train_probe_ids.json").write_text(json.dumps([r["id"] for r in eval_records if r["split"] == "train"], indent=2), encoding="utf-8") dynamic_metrics: List[dict] = [] eval_log_path = args.output_dir / "dynamic_psnr.jsonl" @torch.no_grad() def evaluate_all(step: int, elapsed_sec: float) -> dict: """Generate every train/test image and report split-level PSNR/SSIM.""" per_sample = [] model.eval() step_root = args.output_dir / "eval_images" / f"step_{step:04d}" for idx, rec in enumerate(eval_records): root = Path(rec["_root"]) inp = Image.open(root / rec["input"]).convert("RGB") target = Image.open(root / rec["target"]).convert("RGB") set_seed(args.seed + idx) pred = inferencer.interleave_inference( [inp, rec["prompt"]], think=False, understanding_output=False, cfg_text_scale=4.0, cfg_img_scale=1.5, cfg_interval=[0.4, 1.0], timestep_shift=3.0, num_timesteps=args.eval_steps, cfg_renorm_min=0.0, cfg_renorm_type="global", )[-1] pred_path = step_root / rec["split"] / f'{rec["id"]}.png' pred_path.parent.mkdir(parents=True, exist_ok=True) pred.save(pred_path) rgb = np.asarray(pred.convert("RGB"), dtype=np.float32) / 255.0 per_sample.append({ "id": rec["id"], "scene": rec["scene"], "split": rec["split"], "path": str(pred_path), "color_strength": float(np.mean(np.std(rgb, axis=2))), **image_metrics(pred, target), }) row = {"step": step, "elapsed_sec": elapsed_sec, "eval_steps": args.eval_steps} for split in ("train", "test"): subset = [x for x in per_sample if x["split"] == split] psnr = np.asarray([x["psnr_db"] for x in subset], dtype=np.float64) ssim = np.asarray([x["ssim"] for x in subset], dtype=np.float64) row[split] = { "count": len(subset), "psnr_mean_db": float(psnr.mean()), "psnr_std_db": float(psnr.std()), "ssim_mean": float(ssim.mean()), "ssim_std": float(ssim.std()), "color_strength": float(np.mean([x["color_strength"] for x in subset])), } row["per_sample"] = per_sample dynamic_metrics.append(row) with eval_log_path.open("a", encoding="utf-8") as f: f.write(json.dumps(row) + "\n") print( f'DYNAMIC_EVAL step={step} train_psnr={row["train"]["psnr_mean_db"]:.6f} ' f'test_psnr={row["test"]["psnr_mean_db"]:.6f} ' f'train_ssim={row["train"]["ssim_mean"]:.6f} test_ssim={row["test"]["ssim_mean"]:.6f}', flush=True, ) model.train() return row started = time.time() eval_log_path.write_text("", encoding="utf-8") initial_eval = evaluate_all(0, time.time() - started) val_before = fixed_train_flow_mse() print(f"FIXED_TRAIN_FLOW_MSE before={val_before:.8f}", flush=True) apply_activation_checkpointing( model, checkpoint_wrapper_fn=functools.partial( checkpoint_wrapper, checkpoint_impl=CheckpointImpl.NO_REENTRANT ), check_fn=lambda module: isinstance(module, Qwen2MoTDecoderLayer), ) optimizer = torch.optim.AdamW(trainable, lr=args.lr, betas=(0.9, 0.95), eps=1e-8, weight_decay=0.0) log_path = args.output_dir / "train_log.jsonl" losses: List[float] = [] model.train() # required for Brainmu's packed-loss routing; LoRA dropout is zero. with log_path.open("w", encoding="utf-8") as logf: for step in range(args.steps): optimizer.zero_grad(set_to_none=True) data = prepared_batches[step % len(prepared_batches)] loss, flow_value, gray_value = batch_loss(data, args.seed + 100_000 + step, backward=True) loss.backward() grad_norm = torch.nn.utils.clip_grad_norm_(trainable, 1.0) optimizer.step() value = float(loss.detach().item()) losses.append(value) row = { "step": step, "mse": value, "flow_mse": float(flow_value), "gray_l1": float(gray_value), "grad_norm": float(grad_norm), "lr": optimizer.param_groups[0]["lr"], "elapsed_sec": time.time() - started, "max_memory_mib": torch.cuda.max_memory_allocated() / 1024**2, } logf.write(json.dumps(row) + "\n") logf.flush() if step % 5 == 0 or step + 1 == args.steps: print(f"TRAIN step={step:04d} total={value:.8f} flow={float(flow_value):.8f} gray_l1={float(gray_value):.8f} grad={float(grad_norm):.5f}", flush=True) completed = step + 1 if completed % args.eval_every == 0 or completed == args.steps: evaluate_all(completed, time.time() - started) save_adapter(adapters, args.output_dir / "checkpoints" / f"adapter_step_{completed:04d}.safetensors") val_after = fixed_train_flow_mse() adapter_path = args.output_dir / "adapter.safetensors" save_adapter(adapters, adapter_path) trained_rows = [x for x in dynamic_metrics if x["step"] > 0] best_row = max(trained_rows, key=lambda x: x["test"]["psnr_mean_db"]) best_checkpoint = args.output_dir / "checkpoints" / f'adapter_step_{best_row["step"]:04d}.safetensors' best_adapter_path = args.output_dir / "adapter_best_test_psnr.safetensors" shutil.copyfile(best_checkpoint, best_adapter_path) first_n = losses[: min(20, len(losses))] last_n = losses[-min(20, len(losses)) :] summary = { "steps": args.steps, "lr": args.lr, "flow_weight": args.flow_weight, "gray_weight": args.gray_weight, "prompt": train_records[0]["prompt"], "train_mse_first_window": float(np.mean(first_n)), "train_mse_last_window": float(np.mean(last_n)), "num_train": len(train_records), "num_test": len(test_records), "dynamic_train_probe_count": args.train_eval_count, "dynamic_test_count": len(selected_test_records), "split_rule": "REDS_orig: 2000 sequence-disjoint train pairs; held-out 200 validation pairs", "fixed_train_flow_mse_before": val_before, "fixed_train_flow_mse_after": val_after, "fixed_train_flow_mse_improvement_ratio": val_before / max(val_after, 1e-12), "dynamic_metrics": dynamic_metrics, "initial_metrics": initial_eval, "final_metrics": dynamic_metrics[-1], "best_test_psnr_step": best_row["step"], "best_test_metrics": best_row, "adapter": str(adapter_path), "best_adapter": str(best_adapter_path), "elapsed_sec": time.time() - started, "max_memory_mib": torch.cuda.max_memory_allocated() / 1024**2, } (args.output_dir / "summary.json").write_text(json.dumps(summary, indent=2), encoding="utf-8") # Metric curves and a small qualitative sheet (two train + two test). import matplotlib.pyplot as plt steps = [x["step"] for x in dynamic_metrics] fig, ax = plt.subplots(figsize=(9, 5)) ax.plot(steps, [x["train"]["psnr_mean_db"] for x in dynamic_metrics], "o-", label="train probe PSNR") ax.plot(steps, [x["test"]["psnr_mean_db"] for x in dynamic_metrics], "o-", label="validation PSNR") ax.set(xlabel="optimizer step", ylabel="mean PSNR (dB)", title="Brainmu generation-expert LoRA") ax.grid(alpha=0.3); ax.legend(); fig.tight_layout() fig.savefig(args.output_dir / "dynamic_psnr.png", dpi=160); plt.close(fig) rows = [] chosen = [r for r in eval_records if r["split"] == "train"][:2] + [r for r in eval_records if r["split"] == "test"][:2] final_step = dynamic_metrics[-1]["step"] for rec in chosen: base_path = args.output_dir / "eval_images" / "step_0000" / rec["split"] / f'{rec["id"]}.png' final_path = args.output_dir / "eval_images" / f"step_{final_step:04d}" / rec["split"] / f'{rec["id"]}.png' imgs = [ Image.open(Path(rec["_root"]) / rec["input"]).convert("RGB"), Image.open(base_path).convert("RGB"), Image.open(final_path).convert("RGB"), Image.open(Path(rec["_root"]) / rec["target"]).convert("RGB"), ] imgs = [im.resize((300, 192)) for im in imgs] row = Image.new("RGB", (1200, 220), "white") draw = ImageDraw.Draw(row) for j, (label, im) in enumerate(zip([f'{rec["split"]} BaseNet input', "step 0", f"step {final_step}", "clean GT"], imgs)): row.paste(im, (j * 300, 28)) draw.text((j * 300 + 4, 5), label, fill="black") rows.append(row) sheet = Image.new("RGB", (1200, 220 * len(rows)), "white") for i, row in enumerate(rows): sheet.paste(row, (0, i * 220)) sheet.save(args.output_dir / "comparison.png") print("TRAINING_COMPLETE " + json.dumps(summary), flush=True) if __name__ == "__main__": main()