anima-polish-checkpoints / scripts /72_merge_checkpoints.py
advokat's picture
Upload scripts/72_merge_checkpoints.py with huggingface_hub
ebba223 verified
Raw
History Blame Contribute Delete
18.5 kB
#!/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()