| |
| """ |
| 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 |
|
|
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| _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": |
| |
| |
| |
| |
| sched = {} |
| for i in range(_NUM_BLOCKS): |
| if i < 10: |
| sched[i] = 0.6 |
| elif i < 20: |
| sched[i] = 1.0 |
| else: |
| sched[i] = 0.4 |
| return sched |
|
|
| if spec == "personality_only": |
| |
| 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": |
| |
| 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": |
| |
| sched = {i: 1.0 - 0.7 * (i / (_NUM_BLOCKS - 1)) for i in range(_NUM_BLOCKS)} |
| return sched |
|
|
| if spec.startswith("custom:"): |
| |
| 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) |
|
|
|
|
| |
| |
| |
| 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) |
|
|
|
|
| |
| |
| |
| 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 |
|
|
| |
| 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 |
|
|
| |
| 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) |
|
|
|
|
|
|
| |
| |
| |
| 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 |
|
|
| |
| 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 |
|
|
|
|
| |
| |
| |
| 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) |
|
|
|
|
| |
| |
| |
| 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() |
|
|