#!/usr/bin/env python3 """ 72_merge_checkpoints.py — Task-arithmetic / TIES / DARE merging of diffusion checkpoints, designed for Anima. Use case: anima-base-v1.0 just released. We have step4739 (our FT off preview3). We want to combine v1.0's improved high-res handling + fidelity with the character/composition learnings from our FT. Construct (Ilharco et al., 2022): task_vector = our_FT − preview3_base # what we trained, in weight space new_model = v1.0_base + α · task_vector Plain task arithmetic (--mode add) does the simple weighted sum. TIES (Yadav et al. 2023) resolves sign conflicts between our delta and v1's drift from preview3; DARE (Yu et al. 2023) drops + rescales delta entries to free up signal headroom; --mode hybrid combines TIES + DARE which is the SOTA for non-linear compositional merges. Usage ----- Plain task arithmetic (start here, fastest): python 72_merge_checkpoints.py \\ --base preview3_base.safetensors \\ --ft step4739.safetensors \\ --target v1.0.safetensors \\ --output anima_v1_step4739_arith.safetensors \\ --mode add --alpha 0.5 TIES (sign-aware): python 72_merge_checkpoints.py ... --mode ties --alpha 0.7 --ties-density 0.7 DARE+TIES (recommended once arith is verified): python 72_merge_checkpoints.py ... --mode hybrid --alpha 0.7 \\ --dare-drop-p 0.5 --ties-density 0.5 """ from __future__ import annotations import argparse import re import sys import time from pathlib import Path import torch from safetensors.torch import load_file, save_file # ---------------------------------------------------------------------- # Per-block α schedule # ---------------------------------------------------------------------- # Cosmos-Predict2-2B has 28 DiT transformer blocks (indices 0..27). # Block index correlates roughly with denoising frequency band: # - Early (0-9) : compositional layout, dominant at high σ steps # - Mid (10-19) : character / anatomy / identity (carries "personality") # - Late (20-27) : fine textures / details, dominant at low σ steps # # A layer-wise schedule lets us apply MORE FT delta where we want our # fine-tune's contribution to win, and LESS where we'd rather keep # v1.0's improvements (typically: fidelity in late layers). _NUM_BLOCKS = 28 def parse_alpha_schedule(spec: str, base_alpha: float) -> dict[int, float]: """Return {block_idx: α_multiplier} for one of the named schedules. The returned multiplier is meant to scale `base_alpha` per-block. Non-block keys (embedders, output head) get base_alpha exactly. """ if not spec or spec == "uniform": return {i: 1.0 for i in range(_NUM_BLOCKS)} if spec == "preserve_personality": # Mid-peak schedule. Mid blocks (carrying our FT character) get # full α; late blocks (where AI gloss creeps in) get reduced α # so v1.0's fidelity dominates; early blocks (composition) get # moderate α — let v1.0 shape composition somewhat. sched = {} for i in range(_NUM_BLOCKS): if i < 10: sched[i] = 0.6 # early: take some FT, lean v1.0 elif i < 20: sched[i] = 1.0 # mid: full FT (personality) else: sched[i] = 0.4 # late: reduce FT (let v1.0 fidelity through) return sched if spec == "personality_only": # Strict: only apply FT delta to mid-block personality region. sched = {i: 0.0 for i in range(_NUM_BLOCKS)} for i in range(10, 20): sched[i] = 1.0 return sched if spec == "late_soft": # FT dominates first 70% of blocks, late detail layers stay v1.0 sched = {} for i in range(_NUM_BLOCKS): sched[i] = 1.0 if i < int(0.7 * _NUM_BLOCKS) else 0.3 return sched if spec == "linear_down": # Smooth taper from full α at block 0 to 0.3 at block 27 sched = {i: 1.0 - 0.7 * (i / (_NUM_BLOCKS - 1)) for i in range(_NUM_BLOCKS)} return sched if spec.startswith("custom:"): # custom:1.0,1.0,...,0.3 — comma-separated 28 values vals = [float(x) for x in spec[len("custom:"):].split(",")] if len(vals) != _NUM_BLOCKS: sys.exit(f"custom schedule needs {_NUM_BLOCKS} values, got {len(vals)}") return {i: vals[i] for i in range(_NUM_BLOCKS)} sys.exit(f"unknown schedule: {spec}") def alpha_for_key(key: str, base_alpha: float, schedule: dict[int, float]) -> float: """Look up the per-block α multiplier for a parameter key. Non-block keys (input embedders, output proj, modulation) get the base_alpha unchanged — they're typically very small and we don't want to zero them out. """ m = re.search(r"blocks\.(\d+)\.", key) if m is None: return base_alpha block_idx = int(m.group(1)) return base_alpha * schedule.get(block_idx, 1.0) # ---------------------------------------------------------------------- # Loading / saving # ---------------------------------------------------------------------- def load_sd(path: Path) -> dict: print(f" [load] {path.name}", flush=True) t0 = time.time() sd = load_file(str(path)) print(f" {len(sd)} keys, {sum(t.numel() for t in sd.values())/1e9:.2f}B params, " f"loaded in {time.time()-t0:.1f}s", flush=True) return sd def save_sd(sd: dict, path: Path, dtype: str = "bfloat16"): target_dt = {"bfloat16": torch.bfloat16, "float16": torch.float16, "float32": torch.float32}[dtype] casted = {k: v.to(target_dt).contiguous() for k, v in sd.items()} print(f" [save] {path.name} ({dtype})", flush=True) t0 = time.time() save_file(casted, str(path)) print(f" wrote in {time.time()-t0:.1f}s " f"({path.stat().st_size/1e9:.2f} GB)", flush=True) # ---------------------------------------------------------------------- # Merge modes # ---------------------------------------------------------------------- def merge_add(target_sd, base_sd, ft_sd, alpha, schedule=None, key_filter=None): """Plain task arithmetic: target + α_k (ft - base) with per-block α_k.""" out = {} skipped = 0 schedule = schedule or {i: 1.0 for i in range(_NUM_BLOCKS)} for k, v_target in target_sd.items(): if key_filter is not None and not key_filter(k): out[k] = v_target continue if k not in base_sd or k not in ft_sd: out[k] = v_target skipped += 1 continue if base_sd[k].shape != v_target.shape or ft_sd[k].shape != v_target.shape: out[k] = v_target skipped += 1 continue a = alpha_for_key(k, alpha, schedule) delta = ft_sd[k].float() - base_sd[k].float() out[k] = (v_target.float() + a * delta).to(v_target.dtype) print(f" [add] base α={alpha} merged {len(out)-skipped}/{len(out)} keys, " f"skipped {skipped}", flush=True) return out def _ties_resolve(delta_ft, delta_target, density): """TIES per-tensor sign-conflict resolution. 1. Trim: keep only top-`density` magnitude entries of delta_ft (zero rest). 2. Sign election: pick the sign with larger total magnitude across delta_ft / delta_target. 3. Disjoint merge: take entries from delta_ft whose sign matches the elected sign. """ abs_ft = delta_ft.abs() if density < 1.0: flat = abs_ft.flatten() k = max(1, int(density * flat.numel())) threshold = torch.kthvalue(flat, flat.numel() - k + 1).values mask_topk = abs_ft >= threshold delta_ft_trimmed = delta_ft * mask_topk else: delta_ft_trimmed = delta_ft # Elect sign: positive vs negative magnitude per parameter (over both deltas) pos_mass = (delta_ft_trimmed.clamp(min=0).abs().sum() + delta_target.clamp(min=0).abs().sum()) neg_mass = (delta_ft_trimmed.clamp(max=0).abs().sum() + delta_target.clamp(max=0).abs().sum()) sign_elect = 1.0 if pos_mass >= neg_mass else -1.0 # Keep entries whose sign matches the elected sign mask_sign = ((delta_ft_trimmed > 0) == (sign_elect > 0)) return delta_ft_trimmed * mask_sign.float() def merge_ties(target_sd, base_sd, ft_sd, alpha, density, key_filter=None): """TIES: trim, sign-elect, then add the resolved delta to target.""" out = {} skipped = 0 for k, v_target in target_sd.items(): if key_filter is not None and not key_filter(k): out[k] = v_target continue if k not in base_sd or k not in ft_sd: out[k] = v_target skipped += 1 continue if base_sd[k].shape != v_target.shape or ft_sd[k].shape != v_target.shape: out[k] = v_target skipped += 1 continue delta_ft = ft_sd[k].float() - base_sd[k].float() delta_target = v_target.float() - base_sd[k].float() resolved = _ties_resolve(delta_ft, delta_target, density) out[k] = (v_target.float() + alpha * resolved).to(v_target.dtype) print(f" [ties] α={alpha} density={density} merged " f"{len(out)-skipped}/{len(out)} keys, skipped {skipped}", flush=True) return out def _dare_drop_rescale(delta, drop_p): """DARE: randomly drop drop_p fraction of delta entries, rescale rest by 1/(1-drop_p).""" if drop_p <= 0: return delta keep_mask = (torch.rand_like(delta) > drop_p).float() return (delta * keep_mask) / max(1.0 - drop_p, 1e-3) # ---------------------------------------------------------------------- # Key-pattern filtering — restrict merge to a subset of weights # ---------------------------------------------------------------------- def make_key_filter(include_re: str | None, exclude_re: str | None): inc = re.compile(include_re) if include_re else None exc = re.compile(exclude_re) if exclude_re else None def _ok(k: str) -> bool: if inc is not None and not inc.search(k): return False if exc is not None and exc.search(k): return False return True return _ok def merge_mag_aware(target_sd, base_sd, ft_sd, alpha, sharpness=1.0, key_filter=None): """Per-element magnitude-aware blending. Where the target (v1) has drifted a lot from base (preview3), trust the target's direction more — assume v1's authors deliberately moved those weights. Where v1 hasn't changed, apply our delta more fully. Per-element novelty weight: novelty = 1 / (1 + sharpness · |delta_target| / σ_target) where σ_target is the per-tensor RMS of delta_target. So elements far above the typical v1-drift get heavily penalised, elements at or below average get the full α applied. """ out = {} skipped = 0 for k, v_target in target_sd.items(): if key_filter is not None and not key_filter(k): out[k] = v_target continue if k not in base_sd or k not in ft_sd: out[k] = v_target skipped += 1 continue if base_sd[k].shape != v_target.shape or ft_sd[k].shape != v_target.shape: out[k] = v_target skipped += 1 continue b = base_sd[k].float() delta_ft = ft_sd[k].float() - b delta_target = v_target.float() - b # Per-tensor RMS of delta_target — the "typical" v1 movement scale sigma_target = max(delta_target.pow(2).mean().sqrt().item(), 1e-8) novelty = 1.0 / (1.0 + sharpness * delta_target.abs() / sigma_target) out[k] = (v_target.float() + alpha * novelty * delta_ft).to(v_target.dtype) print(f" [mag_aware] α={alpha} sharpness={sharpness} merged " f"{len(out)-skipped}/{len(out)} keys, skipped {skipped}", flush=True) return out def merge_hybrid(target_sd, base_sd, ft_sd, alpha, density, drop_p, schedule=None, key_filter=None): """DARE drop on delta_ft, then TIES sign-elect, then add to target with per-block α.""" out = {} skipped = 0 schedule = schedule or {i: 1.0 for i in range(_NUM_BLOCKS)} for k, v_target in target_sd.items(): if key_filter is not None and not key_filter(k): out[k] = v_target continue if k not in base_sd or k not in ft_sd: out[k] = v_target skipped += 1 continue if base_sd[k].shape != v_target.shape or ft_sd[k].shape != v_target.shape: out[k] = v_target skipped += 1 continue a = alpha_for_key(k, alpha, schedule) delta_ft = ft_sd[k].float() - base_sd[k].float() delta_target = v_target.float() - base_sd[k].float() delta_ft_dare = _dare_drop_rescale(delta_ft, drop_p) resolved = _ties_resolve(delta_ft_dare, delta_target, density) out[k] = (v_target.float() + a * resolved).to(v_target.dtype) print(f" [hybrid DARE+TIES] base α={alpha} density={density} drop_p={drop_p} " f"merged {len(out)-skipped}/{len(out)} keys, skipped {skipped}", flush=True) return out # ---------------------------------------------------------------------- # Diagnostics # ---------------------------------------------------------------------- def report_drift(base_sd, ft_sd, target_sd): """Print how far each pair of checkpoints has drifted in weight space. Useful sanity check before merging — if v1.0 has drifted very far from preview3, the task vector might point in a direction that no longer makes sense in v1.0's parameter space. """ total_ft = total_target = total_n = 0 for k in base_sd: if k not in ft_sd or k not in target_sd: continue if base_sd[k].shape != ft_sd[k].shape or base_sd[k].shape != target_sd[k].shape: continue b = base_sd[k].float() total_ft += (ft_sd[k].float() - b).pow(2).sum().item() total_target += (target_sd[k].float() - b).pow(2).sum().item() total_n += b.numel() rms_ft = (total_ft / total_n) ** 0.5 rms_target = (total_target / total_n) ** 0.5 print(f" [drift] vs preview3: FT RMS = {rms_ft:.5f} | " f"v1.0 RMS = {rms_target:.5f} | ratio (v1/FT) = {rms_target/rms_ft:.2f}", flush=True) # ---------------------------------------------------------------------- # CLI # ---------------------------------------------------------------------- def main(): p = argparse.ArgumentParser(description=__doc__.split("\n\n")[0]) p.add_argument("--base", type=Path, required=True, help="Reference base checkpoint (preview3) — what FT diverged FROM.") p.add_argument("--ft", type=Path, required=True, help="Fine-tuned checkpoint (step4739).") p.add_argument("--target", type=Path, required=True, help="Target base to merge INTO (anima-base-v1.0).") p.add_argument("--output", type=Path, required=True) p.add_argument("--mode", choices=["add", "ties", "hybrid", "mag_aware"], default="add") p.add_argument("--mag-sharpness", type=float, default=1.0, help="Sharpness of the novelty curve (mag_aware only). " "Higher = more aggressive penalty for elements where " "target already moved heavily.") p.add_argument("--alpha-schedule", default="uniform", help="Per-block α multiplier schedule. Options: " "uniform (default), preserve_personality (mid α=1.0, " "early=0.6, late=0.4), personality_only (α=0 outside " "blocks 10-19), late_soft (α=1.0 first 70%%, 0.3 rest), " "linear_down (taper 1.0→0.3), or custom:v0,v1,...,v27") p.add_argument("--alpha", type=float, default=0.5, help="Strength of the task vector applied to target. " "0=target unchanged; 1=full FT delta added.") p.add_argument("--ties-density", type=float, default=0.7, help="Top-k fraction of delta_ft entries kept by magnitude (TIES/hybrid).") p.add_argument("--dare-drop-p", type=float, default=0.5, help="Random drop probability for delta_ft entries (hybrid only).") p.add_argument("--key-include-pattern", type=str, default=None, help="Regex: only keys matching this are merged. Others kept as target.") p.add_argument("--key-exclude-pattern", type=str, default=None, help="Regex: keys matching this are skipped (kept as target).") p.add_argument("--save-dtype", default="bfloat16") p.add_argument("--no-report", action="store_true") args = p.parse_args() print("[merge] loading...", flush=True) base_sd = load_sd(args.base) ft_sd = load_sd(args.ft) target_sd = load_sd(args.target) if not args.no_report: report_drift(base_sd, ft_sd, target_sd) schedule = parse_alpha_schedule(args.alpha_schedule, args.alpha) if args.alpha_schedule != "uniform": sched_str = ", ".join(f"{i}:{schedule[i]:.2f}" for i in [0, 9, 10, 19, 20, 27]) print(f" [schedule] {args.alpha_schedule} (samples — {sched_str})", flush=True) key_filter = make_key_filter(args.key_include_pattern, args.key_exclude_pattern) if args.key_include_pattern or args.key_exclude_pattern: kept = sum(1 for k in target_sd if key_filter(k)) print(f" [key-filter] include={args.key_include_pattern!r} exclude={args.key_exclude_pattern!r} -> {kept}/{len(target_sd)} keys eligible to merge", flush=True) print(f"\n[merge] mode={args.mode}", flush=True) if args.mode == "add": out = merge_add(target_sd, base_sd, ft_sd, args.alpha, schedule, key_filter=key_filter) elif args.mode == "ties": out = merge_ties(target_sd, base_sd, ft_sd, args.alpha, args.ties_density, key_filter=key_filter) elif args.mode == "mag_aware": out = merge_mag_aware(target_sd, base_sd, ft_sd, args.alpha, args.mag_sharpness, key_filter=key_filter) else: out = merge_hybrid(target_sd, base_sd, ft_sd, args.alpha, args.ties_density, args.dare_drop_p, schedule, key_filter=key_filter) save_sd(out, args.output, dtype=args.save_dtype) print(f"\n[done] {args.output}", flush=True) if __name__ == "__main__": main()