Brainmu-SpikeCamera / src /code /train_brainmu_lora.py
sunbaby's picture
Upload 69 files
4719196
Raw History Blame Contribute Delete
28 kB
#!/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()