Instructions to use BAAI/Brainmu-SpikeCamera with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use BAAI/Brainmu-SpikeCamera with Transformers:
# pip install -U transformers accelerate # Load model directly from transformers import SpikeConvFrontend model = SpikeConvFrontend.from_pretrained("BAAI/Brainmu-SpikeCamera", device_map="auto") - Notebooks
- Google Colab
- Kaggle
Download src/code/train_brainmu_lora.py from BAAI/Brainmu-SpikeCamera: direct link, hf CLI and curl.
- Browser
- Download file 28 kB
-
https://huggingface.co/BAAI/Brainmu-SpikeCamera/resolve/main/src/code/train_brainmu_lora.py
- Command line
-
hf download hf://BAAI/Brainmu-SpikeCamera/src/code/train_brainmu_lora.py
-
curl -L -o train_brainmu_lora.py https://huggingface.co/BAAI/Brainmu-SpikeCamera/resolve/main/src/code/train_brainmu_lora.py
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, | |
| ) | |
| 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] | |
| 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" | |
| 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() | |