#!/usr/bin/env python3 """graft_krea_to_zimage.py Cross-model weight graft from Krea 2 into Zimage (z_image_turbo). This is a reference implementation of the cross-model graft technique demonstrated on a real pair of pretrained image diffusion transformers. It transfers Krea 2's learned attention head specializations and MLP feature transforms into Zimage's parameter slots, respecting per-head geometry and handling the dimensional reduction from Krea's larger tensors to Zimage's smaller ones. No training. No dataset. No gradient descent. Just directed substitution of weight subspaces in architecturally matched slots. ============================================================================== THE GRAFT MECHANIC (universal to this technique) ============================================================================== Given two transformer models with: - the same head_dim (128 in both here) - the same MLP family (both SwiGLU here) - compatible attention layouts (both fused QKV here) ...you can transfer donor character into a target by: 1. Iterate over blocks; for each target block choose a donor block to source from (a "block map" — shift, peak-centered, or cycled). 2. For each attention head slot in the target, choose which donor head to place there ("source-head-selection" + "target-head-selection"). 3. Extract that donor head as a small tensor slice, reduce its dimensions to fit the target slot (truncation or SVD-based projection). 4. Blend the reduced donor slice into the target slice using bounded rotation: - `linear` : direct addition - `residual` : add the delta between donor and target - `linear-mag` : add + rescale to preserve target's magnitude (this is the safe default at higher strengths — strength becomes rotation angle rather than amplification) 5. Optionally strip the component of donor that aligns with target ("orthogonal projection") so the graft only adds directions the target wasn't using. 6. Repeat for out_proj (per-head column mixing) and MLP tensors (feature-space transformations). The result: target attention heads and MLPs rotate toward donor's directions while preserving target's architectural roles. Character transfers. ============================================================================== KREA 2 -> ZIMAGE SPECIFICS ============================================================================== Zimage architecture (target): layers.N.attention.qkv.weight (11520, 3840) fused Q|K|V, 30h * 128 layers.N.attention.out.weight (3840, 3840) layers.N.attention.q_norm.weight (128,) (not grafted) layers.N.attention.k_norm.weight (128,) (not grafted) layers.N.feed_forward.w1.weight (10240, 3840) SwiGLU value (up) layers.N.feed_forward.w3.weight (10240, 3840) SwiGLU gate layers.N.feed_forward.w2.weight (3840, 10240) down projection (context_refiner.N.* uses identical structure; 2 blocks vs 30 main layers) Krea 2 architecture (donor): .blocks.N.attn.wq.weight (6144, 6144) 48 heads * 128 .blocks.N.attn.wk.weight (1536, 6144) 12 heads * 128 (GQA 4:1) .blocks.N.attn.wv.weight (1536, 6144) 12 heads * 128 (GQA 4:1) .blocks.N.attn.wo.weight (6144, 6144) .blocks.N.mlp.up.weight (16384, 6144) SwiGLU value .blocks.N.mlp.gate.weight (16384, 6144) SwiGLU gate .blocks.N.mlp.down.weight (6144, 16384) Coverage per grafted block: Q: 30 target slots draw from 30 of Krea's 48 Q heads (63% selection) x 3840/6144 = 62% column coverage per head K/V: 12 Krea heads -> first 12 of 30 target slots (40% slot use) x 3840/6144 = 62% column coverage per head MLP w1/w3: 10240/16384 rows x 3840/6144 cols = 62% x 62% coverage MLP w2: 3840/6144 rows x 10240/16384 cols = 62% x 62% coverage ============================================================================== WHY IT WORKS (short version) ============================================================================== Multi-head attention was designed so heads specialize on distinct patterns. A head's "role" is fixed by its slot position in the network; a head's "specialization" lives in its weights. Substituting one model's specialized head for another's — at the same slot — transfers the specialization while preserving the role. Character in image models lives in these specializations more than in specific parameter values. Two heads with the same magnitude but different singular directions produce completely different outputs on the same input. So `linear-mag` mode is the mathematically clean way to preserve target's operational scale while rotating its direction toward donor's direction: bounded rotation, magnitude identity, character transfers. """ import os import math import sys import argparse from typing import Dict, List, Optional, Tuple import torch from safetensors import safe_open from safetensors.torch import save_file from tqdm import tqdm # ============================================================================ # Constants — architectural dimensions of Krea 2 (donor) and Zimage (target) # ============================================================================ # Donor: Krea 2 KREA_TOTAL_BLOCKS = 28 KREA_HIDDEN = 6144 KREA_Q_HEADS = 48 KREA_KV_HEADS = 12 # Grouped-query attention: 4 Q heads per K/V head KREA_HEAD_DIM = 128 KREA_Q_INNER = KREA_Q_HEADS * KREA_HEAD_DIM # 6144 KREA_KV_INNER = KREA_KV_HEADS * KREA_HEAD_DIM # 1536 KREA_MLP_INNER = 16384 KREA_PREFIX_CANDIDATES = [ "model.diffusion_model.blocks", "blocks", ] # Target: Zimage ZIMAGE_MAIN_BLOCKS = 30 ZIMAGE_REFINER_BLOCKS = 2 ZIMAGE_HIDDEN = 3840 ZIMAGE_HEADS = 30 ZIMAGE_HEAD_DIM = 128 ZIMAGE_INNER = ZIMAGE_HEADS * ZIMAGE_HEAD_DIM # 3840 ZIMAGE_MLP_INNER = 10240 ZIMAGE_MAIN_PREFIX = "layers" ZIMAGE_REFINER_PREFIX = "context_refiner" EPS = 1e-8 # ============================================================================ # Dimensional reduction: truncation vs SVD projection # ---------------------------------------------------------------------------- # Donor tensors are larger than target slots in every dimension. Two ways to # reduce them: # # truncate : drop trailing rows or columns (simple, fast, loses everything # in the dropped indices). # # svd : project onto the top-N singular directions (preserves the # "most important" directions the donor tensor spans). # # SVD only helps when the target reduction dim is <= the matrix rank. # For per-head slices (128 x hidden_dim), rank is capped at 128, so requesting # more than 128 dims via SVD falls back to truncation automatically. # ============================================================================ def reduce_dim_via_svd(t: torch.Tensor, target_dim: int, axis: int, device: torch.device) -> torch.Tensor: """Reduce t along `axis` to `target_dim` via truncated SVD projection. axis=0 reduces rows, axis=1 reduces columns. When target_dim > rank(t) = min(rows, cols), SVD literally cannot span that many independent directions — falls back to truncation. """ orig_dtype = t.dtype tf = t.to(device=device, dtype=torch.float32) rows, cols = tf.shape r_max = min(rows, cols) # Rank ceiling: SVD adds nothing beyond truncation when target > rank if target_dim > r_max: if axis == 1: return tf[:, :target_dim].to(dtype=orig_dtype).cpu().contiguous() else: return tf[:target_dim, :].to(dtype=orig_dtype).cpu().contiguous() # SVD path — target_dim fits inside the achievable rank # driver='gesvd' uses the stable LAPACK routine (avoids the cuSOLVER # SGESVDJ Windows bug on large tensors). CPU fallback for safety. try: U, S, Vh = torch.linalg.svd(tf, full_matrices=False, driver='gesvd') except Exception: U, S, Vh = torch.linalg.svd(tf.cpu(), full_matrices=False) U, S, Vh = U.to(device), S.to(device), Vh.to(device) if axis == 1: # Rank-target_dim reconstruction in column space # Result shape: (rows, target_dim), preserving top-target_dim directions result = U[:, :target_dim] * S[:target_dim] elif axis == 0: # Rank-target_dim reconstruction in row space # Result shape: (target_dim, cols) result = S[:target_dim].unsqueeze(1) * Vh[:target_dim, :] else: raise ValueError(f"axis must be 0 or 1, got {axis}") return result.to(dtype=orig_dtype).cpu().contiguous() def reduce_dim(t: torch.Tensor, target_dim: int, axis: int, mode: str, device: torch.device) -> torch.Tensor: """Reduce t along `axis` to `target_dim` via mode='truncate' or 'svd'.""" if t.shape[axis] == target_dim: return t.contiguous() if t.shape[axis] < target_dim: raise ValueError(f"Cannot expand dim {axis}: {t.shape[axis]} -> {target_dim}") if mode == "truncate": return (t[:target_dim, :] if axis == 0 else t[:, :target_dim]).contiguous() elif mode == "svd": return reduce_dim_via_svd(t, target_dim, axis, device) else: raise ValueError(f"Unknown reduce mode: {mode!r}") def reduce_2d(t: torch.Tensor, target_rows: int, target_cols: int, mode: str, device: torch.device) -> torch.Tensor: """Reduce t along BOTH dims to (target_rows, target_cols). Applied to MLP tensors where both dimensions differ between donor and target. Two-pass: reduce rows first, then cols. Not mathematically optimal (that would be a Tucker decomposition) but simple and effective. """ r_now, c_now = t.shape if r_now == target_rows and c_now == target_cols: return t.contiguous() if r_now > target_rows: t = reduce_dim(t, target_rows, 0, mode, device) if t.shape[1] > target_cols: t = reduce_dim(t, target_cols, 1, mode, device) return t # ============================================================================ # Head selection: which donor heads go into which target slots # ---------------------------------------------------------------------------- # Krea has 48 Q heads, Zimage has 30. We must choose 30 of Krea's 48. Different # selections produce different character (which heads "specialize" in what # isn't documented for either model, so this is empirical). # ============================================================================ def select_source_q_heads(mode: str) -> List[int]: """Choose which Krea Q heads (of 48) to source from.""" if mode == "first30": return list(range(0, 30)) if mode == "middle30": # Skip 9 heads on each side (structural early, refinement late) return list(range(9, 39)) if mode == "last30": return list(range(18, 48)) if mode == "spread30": # Evenly sample across all 48 heads return [round(i * 47 / 29) for i in range(30)] if mode == "groups-first": # Respect Krea's GQA grouping: 12 K/V groups * 4 Q heads per group. # Take one Q head from each group first, then the next, etc. picks = [] for offset in range(4): for group in range(KREA_KV_HEADS): picks.append(group * 4 + offset) if len(picks) == 30: return picks return picks[:30] if mode == "all48": # All 48 heads returned; if target has fewer slots the caller truncates return list(range(48)) raise ValueError(f"Unknown source-head-selection mode: {mode!r}") def select_target_q_slots(mode: str, n_needed: int) -> List[int]: """Choose which Zimage Q slots (of 30) receive donor content.""" if n_needed > ZIMAGE_HEADS: n_needed = ZIMAGE_HEADS if mode == "first": return list(range(0, n_needed)) if mode == "last": return list(range(ZIMAGE_HEADS - n_needed, ZIMAGE_HEADS)) if mode == "middle": start = (ZIMAGE_HEADS - n_needed) // 2 return list(range(start, start + n_needed)) if mode == "spread": return [round(i * (ZIMAGE_HEADS - 1) / max(1, n_needed - 1)) for i in range(n_needed)] if mode == "all": return list(range(0, min(n_needed, ZIMAGE_HEADS))) raise ValueError(f"Unknown target-head-selection mode: {mode!r}") # ============================================================================ # Block mapping: for each target block, which donor block do we sample from? # ---------------------------------------------------------------------------- # The relationship between donor block position and target block position # matters because early/middle/late blocks in transformers tend to specialize # in different aspects (structure, character, detail respectively). # ============================================================================ def build_block_map(zimage_targets: List[int], mode: str, donor_total: int, target_total: int, donor_shift: int, donor_peak: int, target_peak: int, clamp: bool = True) -> List[Tuple[int, int]]: """Return list of (target_block, donor_block) pairs. Modes: shift : linear 1:1 with a start offset (target 0 -> donor donor_shift) peak : center donor_peak on target_peak, scale proportionally linear : evenly distribute donor blocks across target range off : simple modulo (target N -> donor N % donor_total) """ pairs = [] for tb in zimage_targets: if mode == "shift": db = tb - min(zimage_targets) + donor_shift elif mode == "peak": offset = (tb - target_peak) * donor_total / max(1, target_total) db = int(round(donor_peak + offset)) elif mode == "linear": frac = (tb - min(zimage_targets)) / max(1, (max(zimage_targets) - min(zimage_targets))) db = int(round(frac * (donor_total - 1))) elif mode == "off": db = tb % donor_total else: raise ValueError(f"Unknown remap mode: {mode!r}") if clamp: db = max(0, min(donor_total - 1, db)) pairs.append((tb, db)) return pairs def parse_block_list(s: str) -> List[int]: """Parse 'A:B' as range(A,B) or 'A,B,C' as explicit list.""" if ":" in s: a, b = s.split(":") return list(range(int(a), int(b))) return [int(x) for x in s.split(",")] # ============================================================================ # The heart of the graft: blending one weight patch into another # ---------------------------------------------------------------------------- # Called for every per-head slice, out_proj column region, and MLP tensor. # The name "blend_patch" refers to the standard math sense of a "patch" — # a rectangular slice of a tensor being modified. # # Three modes trade off differently: # # linear : new = base + strength * donor # Simple additive. Magnitude of result drifts up (or down) from # base. At high strengths, this cascades through the network. # # residual : new = base + strength * (donor - base) # Adds the *delta*, so strength=1 replaces base with donor. # Better for iterative composition but still has magnitude drift. # # linear-mag : new = base + strength * donor, then rescale to |base| # Rotates base toward donor while preserving base's magnitude. # Strength becomes a rotation angle (bounded to 90 deg max, at # strength -> infinity). This is the mathematically cleanest # option for cross-model transfer at higher strengths. # # The orthogonal option (recommended default) first strips donor's component # that aligns with base's direction. Then only the perpendicular part gets # blended in — meaning the graft ADDS directions base wasn't using, rather # than overwriting directions base was using. By Pythagorean argument, this # cannot reduce base's response in any direction it was previously responsive # to. # ============================================================================ def blend_patch(h_patch: torch.Tensor, donor: torch.Tensor, strength: float, mode: str, orthogonal: bool, eps: float = EPS) -> Tuple[torch.Tensor, float, float]: """Blend `donor` into `h_patch` with the given strength and mode. Returns (blended_tensor, arc_deg, rel_residual) — the arc is the angle between h_patch and blended (how much the base rotated); rel_residual is the fractional magnitude of the change. """ # Orthogonal projection: remove donor's component along h_patch direction if orthogonal: flat_h = h_patch.flatten() flat_d = donor.flatten() nh = flat_h.norm() + eps proj_coef = (flat_d * flat_h).sum() / (nh * nh) donor_eff = donor - proj_coef * h_patch else: donor_eff = donor # Apply the chosen blend mode if mode == "linear": blended = h_patch + strength * donor_eff elif mode == "residual": diff = donor_eff - h_patch blended = h_patch + strength * diff elif mode == "linear-mag": # Bounded rotation: add + rescale to preserve |h_patch|. # If donor_eff is perpendicular to h_patch, rotation angle = # arctan(strength * |donor_eff| / |h_patch|). tmp = h_patch + strength * donor_eff h_norm = h_patch.norm() + eps tmp_norm = tmp.norm() + eps blended = tmp * (h_norm / tmp_norm) else: raise ValueError(f"Unknown blend mode: {mode!r}") # Diagnostics for logging vh = h_patch.flatten() vb = blended.flatten() dot = torch.clamp((vh * vb).sum() / ((vh.norm() + eps) * (vb.norm() + eps)), -1.0, 1.0) arc = math.degrees(float(torch.acos(dot))) rel = float((blended - h_patch).norm() / (h_patch.norm() + eps)) return blended, arc, rel # ============================================================================ # Per-band graft functions # ---------------------------------------------------------------------------- # Each band (Q, K/V, out_proj, MLP up/gate/down) needs its own placement logic # because the tensor shapes and per-head slicing differ. # ============================================================================ def apply_qkv_graft(zimage_base: Dict[str, torch.Tensor], zimage_prefix: str, krea_wq: torch.Tensor, krea_wk: Optional[torch.Tensor], krea_wv: Optional[torch.Tensor], q_strength: float, kv_strength: float, do_kv: bool, mode: str, orthogonal: bool, source_q_heads: List[int], target_q_slots: List[int], hidden_reduce: str, dev: torch.device) -> Tuple[Dict[str, torch.Tensor], Dict[str, list]]: """Graft Krea Q (and optionally K/V) into Zimage's fused attention.qkv. Zimage stores Q, K, V concatenated along axis 0: rows [0 : ZIMAGE_INNER) = Q band (30 heads) rows [ZIMAGE_INNER : 2*ZIMAGE_INNER) = K band (30 heads) rows [2*ZIMAGE_INNER : 3*ZIMAGE_INNER) = V band (30 heads) Each head is 128 rows (head_dim). To graft head H of the Q band we modify rows [H*128 : (H+1)*128] within the Q band. """ qkv_key = zimage_prefix + "attention.qkv.weight" if qkv_key not in zimage_base: raise KeyError(f"Zimage qkv not found: {qkv_key}") zimage_qkv = zimage_base[qkv_key].to(dev, dtype=torch.float32).clone() stats = {"q_arc": [], "q_rel": [], "kv_arc": [], "kv_rel": []} # === Q graft: per-head placement === # Krea wq has shape (48*128, 6144). View as 48 heads of shape (128, 6144). krea_q_view = krea_wq.view(KREA_Q_HEADS, KREA_HEAD_DIM, KREA_HIDDEN) for source_head, target_slot in zip(source_q_heads, target_q_slots): # Extract one Krea head: (128, 6144) krea_head = krea_q_view[source_head].to(dev, dtype=torch.float32) # Reduce hidden dim: (128, 6144) -> (128, 3840) # For per-head slices, rank is only 128, so this falls back to truncation # regardless of --hidden-reduce mode. krea_head_reduced = reduce_dim(krea_head, ZIMAGE_HIDDEN, 1, hidden_reduce, dev).to(dev) # Place into target slot's rows within the Q band r0 = target_slot * ZIMAGE_HEAD_DIM r1 = r0 + ZIMAGE_HEAD_DIM zimage_slice = zimage_qkv[r0:r1, :].clone() blended, arc, rel = blend_patch(zimage_slice, krea_head_reduced, q_strength, mode, orthogonal) zimage_qkv[r0:r1, :] = blended stats["q_arc"].append(arc) stats["q_rel"].append(rel) # === K/V graft (optional): 12 GQA heads into first 12 target slots === if do_kv and krea_wk is not None and krea_wv is not None: # Krea uses grouped-query attention: only 12 K/V heads exist (vs 48 Q). # Zimage has 30 K/V slots and no GQA. We place Krea's 12 K/V heads # into the first 12 Zimage K/V slots; slots 12-29 stay native. krea_k_view = krea_wk.view(KREA_KV_HEADS, KREA_HEAD_DIM, KREA_HIDDEN) krea_v_view = krea_wv.view(KREA_KV_HEADS, KREA_HEAD_DIM, KREA_HIDDEN) n_kv = min(KREA_KV_HEADS, ZIMAGE_HEADS) for kv_head in range(n_kv): # K band offset: rows [ZIMAGE_INNER + kv_head*128 : ...] k_r0 = ZIMAGE_INNER + kv_head * ZIMAGE_HEAD_DIM k_r1 = k_r0 + ZIMAGE_HEAD_DIM krea_k_head = krea_k_view[kv_head].to(dev, dtype=torch.float32) krea_k_reduced = reduce_dim(krea_k_head, ZIMAGE_HIDDEN, 1, hidden_reduce, dev).to(dev) z_k_slice = zimage_qkv[k_r0:k_r1, :].clone() blended_k, k_arc, k_rel = blend_patch(z_k_slice, krea_k_reduced, kv_strength, mode, orthogonal) zimage_qkv[k_r0:k_r1, :] = blended_k stats["kv_arc"].append(k_arc) stats["kv_rel"].append(k_rel) # V band offset: rows [2*ZIMAGE_INNER + kv_head*128 : ...] v_r0 = 2 * ZIMAGE_INNER + kv_head * ZIMAGE_HEAD_DIM v_r1 = v_r0 + ZIMAGE_HEAD_DIM krea_v_head = krea_v_view[kv_head].to(dev, dtype=torch.float32) krea_v_reduced = reduce_dim(krea_v_head, ZIMAGE_HIDDEN, 1, hidden_reduce, dev).to(dev) z_v_slice = zimage_qkv[v_r0:v_r1, :].clone() blended_v, v_arc, v_rel = blend_patch(z_v_slice, krea_v_reduced, kv_strength, mode, orthogonal) zimage_qkv[v_r0:v_r1, :] = blended_v stats["kv_arc"].append(v_arc) stats["kv_rel"].append(v_rel) return {qkv_key: zimage_qkv.to(torch.bfloat16).cpu()}, stats def apply_out_proj_graft(zimage_base: Dict[str, torch.Tensor], zimage_prefix: str, krea_wo: torch.Tensor, strength: float, mode: str, orthogonal: bool, source_q_heads: List[int], target_q_slots: List[int], hidden_reduce: str, dev: torch.device) -> Tuple[Dict[str, torch.Tensor], list]: """Graft Krea attn.wo into Zimage attention.out. out_proj mixes per-head outputs into the residual stream. Its columns are the per-head contribution regions: cols [h*128 : (h+1)*128] = head h's output contribution to the mix We match this per-head structure: Krea head H contributes to the same target slot chosen for the Q graft. Row dim is the hidden dim (donor 6144 -> target 3840, reduced once globally before per-head col placement). """ out_key = zimage_prefix + "attention.out.weight" if out_key not in zimage_base: raise KeyError(f"Zimage attention.out not found: {out_key}") zimage_out = zimage_base[out_key].to(dev, dtype=torch.float32).clone() arcs = [] # Reduce Krea wo rows 6144 -> 3840 once (shared across all head columns) krea_wo_reduced_rows = reduce_dim(krea_wo, ZIMAGE_HIDDEN, 0, hidden_reduce, dev).to(dev, dtype=torch.float32) # Then place per-head column regions into matching target slots for source_head, target_slot in zip(source_q_heads, target_q_slots): c0k = source_head * KREA_HEAD_DIM c1k = c0k + KREA_HEAD_DIM krea_head_cols = krea_wo_reduced_rows[:, c0k:c1k] # (3840, 128) c0z = target_slot * ZIMAGE_HEAD_DIM c1z = c0z + ZIMAGE_HEAD_DIM z_slice = zimage_out[:, c0z:c1z].clone() blended, arc, _ = blend_patch(z_slice, krea_head_cols, strength, mode, orthogonal) zimage_out[:, c0z:c1z] = blended arcs.append(arc) return {out_key: zimage_out.to(torch.bfloat16).cpu()}, arcs def apply_mlp_graft(zimage_base: Dict[str, torch.Tensor], zimage_prefix: str, krea_up: torch.Tensor, krea_gate: torch.Tensor, krea_down: Optional[torch.Tensor], strength: float, mode: str, orthogonal: bool, include_gate: bool, include_w2: bool, hidden_reduce: str, dev: torch.device) -> Tuple[Dict[str, torch.Tensor], Dict[str, float]]: """Graft Krea SwiGLU MLP into Zimage feed_forward. SwiGLU is a three-tensor MLP: out = down(silu(gate(x)) * up(x)) Both Krea and Zimage use this family. Direct 1:1 name mapping: Krea mlp.up -> Zimage feed_forward.w1 (value) Krea mlp.gate -> Zimage feed_forward.w3 (gate, opt-in default ON) Krea mlp.down -> Zimage feed_forward.w2 (down, opt-in default OFF) All three require 2D reduction since both inner (16384->10240) and hidden (6144->3840) dims differ. Down projection (w2) is left off by default because it tends to be more destructive to grafts than up/gate — it's the "output-shaping" tensor of the MLP and target-native w2 keeps outputs consistent with the rest of the network's expectations. """ out = {} stats = {} # w1 (up / SwiGLU value) w1_key = zimage_prefix + "feed_forward.w1.weight" if w1_key in zimage_base: z_w1 = zimage_base[w1_key].to(dev, dtype=torch.float32).clone() krea_up_reduced = reduce_2d(krea_up, ZIMAGE_MLP_INNER, ZIMAGE_HIDDEN, hidden_reduce, dev).to(dev, dtype=torch.float32) blended, arc, _ = blend_patch(z_w1, krea_up_reduced, strength, mode, orthogonal) out[w1_key] = blended.to(torch.bfloat16).cpu() stats["w1_arc"] = arc # w3 (SwiGLU gate) if include_gate: w3_key = zimage_prefix + "feed_forward.w3.weight" if w3_key in zimage_base: z_w3 = zimage_base[w3_key].to(dev, dtype=torch.float32).clone() krea_gate_reduced = reduce_2d(krea_gate, ZIMAGE_MLP_INNER, ZIMAGE_HIDDEN, hidden_reduce, dev).to(dev, dtype=torch.float32) blended, arc, _ = blend_patch(z_w3, krea_gate_reduced, strength, mode, orthogonal) out[w3_key] = blended.to(torch.bfloat16).cpu() stats["w3_arc"] = arc # w2 (down projection) if include_w2 and krea_down is not None: w2_key = zimage_prefix + "feed_forward.w2.weight" if w2_key in zimage_base: z_w2 = zimage_base[w2_key].to(dev, dtype=torch.float32).clone() krea_down_reduced = reduce_2d(krea_down, ZIMAGE_HIDDEN, ZIMAGE_MLP_INNER, hidden_reduce, dev).to(dev, dtype=torch.float32) blended, arc, _ = blend_patch(z_w2, krea_down_reduced, strength, mode, orthogonal) out[w2_key] = blended.to(torch.bfloat16).cpu() stats["w2_arc"] = arc return out, stats # ============================================================================ # Main # ============================================================================ def main(): ap = argparse.ArgumentParser( description="Graft Krea 2 weights into Zimage. Cross-model weight " "transfer via per-head placement and dimensional reduction, " "no training involved.") ap.add_argument("--krea-donor", required=True, help="Path to Krea 2 checkpoint (safetensors).") ap.add_argument("--zimage-base", required=True, help="Path to Zimage base checkpoint (safetensors).") ap.add_argument("--output", required=True, help="Path to write the grafted Zimage checkpoint.") ap.add_argument("--target", choices=["main", "refiner"], default="main", help="Which Zimage sub-stack to graft into. " "'main' targets layers.N (30 blocks — the main " "denoising stack, primary character carrier). " "'refiner' targets context_refiner.N (2 blocks — " "text conditioning refinement; small but influential).") ap.add_argument("--blocks", type=str, default=None, help="Target block range (e.g. '0:30' for all main layers, " "'15:30' for late layers only). Defaults to all blocks " "of the chosen target.") ap.add_argument("--strength", type=float, default=0.15, help="Q strength (also V strength if --kv-strength not set). " "Krea grafts run un-gated in some target contexts, so " "effective magnitude can be higher than the value " "suggests. Start moderate (0.10-0.20) and iterate. " "With --mode linear-mag, strength has a bounded " "rotation-angle interpretation and is safer at higher " "values.") ap.add_argument("--mode", choices=["linear", "residual", "linear-mag"], default="residual", help="Blend mode. 'linear': base + s*donor (standard). " "'residual' (default): base + s*(donor - base). " "'linear-mag': add then rescale to preserve base " "magnitude — recommended for higher strengths since " "it's mathematically bounded and prevents amplitude " "cascade through the network.") ap.add_argument("--orthogonal", action="store_true", default=True, help="Strip donor's base-parallel component before blending. " "Default ON. Preserves base capabilities in principal " "directions — the graft can only add directions the " "target wasn't using, not overwrite ones it was.") ap.add_argument("--no-orthogonal", dest="orthogonal", action="store_false") ap.add_argument("--remap", choices=["shift", "peak", "linear", "off"], default="shift", help="How to map target blocks to donor blocks. " "'shift' (default): 1:1 with an offset. " "'peak': center donor peak on target peak, proportional. " "'linear': evenly distribute donor blocks over target range. " "'off': modulo mapping (wraps).") ap.add_argument("--donor-shift", type=int, default=0, help="Offset for --remap shift.") ap.add_argument("--krea-peak", type=int, default=14, help="Krea block index at peak position (for --remap peak).") ap.add_argument("--zimage-peak", type=int, default=None, help="Target block index to align with Krea peak (for --remap peak). " "Defaults to midpoint of the target range.") ap.add_argument("--source-head-selection", choices=["first30", "middle30", "last30", "spread30", "groups-first", "all48"], default="middle30", help="Which 30 of Krea's 48 Q heads to source from. " "'middle30' (default) skips the first and last 9 heads. " "'spread30' samples evenly across all 48. " "'groups-first' respects Krea's GQA groupings. " "'all48' returns 48; if target has fewer slots the " "last 30 are used.") ap.add_argument("--target-head-selection", choices=["first", "last", "middle", "spread", "all"], default="all", help="Which Zimage Q slots receive donor content. " "'all' (default) uses all 30 target slots.") ap.add_argument("--kv", action="store_true", default=False, help="Also graft K/V. Krea's GQA (12 K/V heads) limits " "coverage to 12 of Zimage's 30 K/V slots. Off by " "default — enable to include the attention routing " "and value projections in the graft.") ap.add_argument("--kv-strength", type=float, default=None, help="K/V strength. Defaults to 0.5 * --strength.") ap.add_argument("--out-proj", action="store_true", default=False, help="Also graft attention.out (head-mixing projection).") ap.add_argument("--out-proj-strength", type=float, default=None, help="attention.out strength. Defaults to 0.5 * --strength.") ap.add_argument("--mlp", action="store_true", default=False, help="Also graft feed_forward w1 (SwiGLU value).") ap.add_argument("--mlp-blocks", type=str, default=None, help="MLP block range. Defaults to --blocks. Use to restrict " "MLP graft to a subset — for example, exclude the last " "few blocks that specialize in detail synthesis matching " "the target's VAE decoder.") ap.add_argument("--mlp-strength", type=float, default=None, help="MLP strength. Defaults to --strength.") ap.add_argument("--include-gate", action="store_true", default=True, help="Include SwiGLU gate (Krea mlp.gate -> Zimage w3). " "Default ON — both models are SwiGLU-native so gate " "transfers cleanly.") ap.add_argument("--no-include-gate", dest="include_gate", action="store_false") ap.add_argument("--mlp-include-w2", action="store_true", default=False, help="Include down projection (Krea mlp.down -> Zimage w2). " "Off by default — grafting w2 is typically more " "destructive than up/gate.") ap.add_argument("--hidden-reduce", choices=["truncate", "svd"], default="svd", help="How to reduce Krea's larger dimensions to Zimage sizes. " "'truncate' drops trailing rows/cols (fast, discards " "information beyond the truncation point). " "'svd' (default) projects onto top singular directions " "when the target dim fits inside matrix rank, otherwise " "falls back to truncation.") ap.add_argument("--cycles", type=int, default=1, help="Number of graft passes. Each pass sees the accumulated " "state from the previous one. Higher cycle counts spread " "character over broader donor subspaces at the cost of " "runtime.") ap.add_argument("--device", choices=["cuda", "cpu"], default="cuda") ap.add_argument("--dry-run", action="store_true", help="Print configuration and exit before loading models.") args = ap.parse_args() # Resolve derived defaults if args.kv_strength is None: args.kv_strength = 0.5 * args.strength if args.out_proj_strength is None: args.out_proj_strength = 0.5 * args.strength if args.mlp_strength is None: args.mlp_strength = args.strength # Resolve target prefix and block count if args.target == "main": target_prefix_base = ZIMAGE_MAIN_PREFIX target_total = ZIMAGE_MAIN_BLOCKS else: target_prefix_base = ZIMAGE_REFINER_PREFIX target_total = ZIMAGE_REFINER_BLOCKS # Parse block ranges if args.blocks is None: zimage_targets = list(range(target_total)) else: zimage_targets = parse_block_list(args.blocks) for b in zimage_targets: if b < 0 or b >= target_total: print(f"ERROR: Zimage {args.target} block {b} out of range " f"[0, {target_total})"); sys.exit(1) if args.mlp_blocks is None: mlp_targets = list(zimage_targets) else: mlp_targets = parse_block_list(args.mlp_blocks) for b in mlp_targets: if b < 0 or b >= target_total: print(f"ERROR: --mlp-blocks contains {b} out of range"); sys.exit(1) mlp_set = set(mlp_targets) if args.zimage_peak is None: args.zimage_peak = (min(zimage_targets) + max(zimage_targets)) // 2 # Head selection source_q_heads = select_source_q_heads(args.source_head_selection) n_heads_selected = len(source_q_heads) target_q_slots = select_target_q_slots(args.target_head_selection, n_heads_selected) n_pairs = min(len(source_q_heads), len(target_q_slots)) source_q_heads = source_q_heads[:n_pairs] target_q_slots = target_q_slots[:n_pairs] # Print configuration banner print("=" * 60) print("graft_krea_to_zimage") print("=" * 60) print(f"Donor : {os.path.basename(args.krea_donor)}") print(f"Base : {os.path.basename(args.zimage_base)}") print(f"Output : {os.path.basename(args.output)}") print(f"Target : Zimage.{args.target} ({target_total} blocks)") print(f"Target blocks: {zimage_targets[0]}:{zimage_targets[-1]+1} ({len(zimage_targets)} blocks)") print(f"Head pairs : {n_pairs} (source={args.source_head_selection}, " f"target={args.target_head_selection})") print(f"Krea heads : {source_q_heads[:5]}... -> Zimage slots: {target_q_slots[:5]}...") print(f"Strengths : Q={args.strength} V={args.strength} " f"K/V={args.kv_strength if args.kv else 'skip'} " f"out_proj={args.out_proj_strength if args.out_proj else 'skip'} " f"MLP={args.mlp_strength if args.mlp else 'skip'}") print(f"Mode : {args.mode} Orthogonal: {args.orthogonal}") print(f"Remap : {args.remap} donor-shift={args.donor_shift} " f"krea-peak={args.krea_peak} zimage-peak={args.zimage_peak}") print(f"Hidden reduce: {args.hidden_reduce}") print(f"Cycles : {args.cycles}") print(f"Device : {args.device} Dry-run: {args.dry_run}") print("=" * 60) if args.dry_run: print("Dry-run: exiting before loading models.") return dev = torch.device(args.device if torch.cuda.is_available() else "cpu") # Load Zimage base print("\nLoading Zimage base...") zimage_base: Dict[str, torch.Tensor] = {} zimage_metadata: Dict[str, str] = {} with safe_open(args.zimage_base, framework="pt", device="cpu") as f: md = f.metadata() if md: zimage_metadata = dict(md) for k in f.keys(): zimage_base[k] = f.get_tensor(k) print(f" {len(zimage_base)} tensors loaded.") # Verify Krea donor and detect its key prefix print("\nDetecting Krea prefix...") fkrea = safe_open(args.krea_donor, framework="pt", device="cpu") krea_keys = set(fkrea.keys()) krea_prefix = None for candidate in KREA_PREFIX_CANDIDATES: if f"{candidate}.0.attn.wq.weight" in krea_keys: krea_prefix = candidate break if krea_prefix is None: print(f"ERROR: could not detect Krea prefix. Tried: " f"{KREA_PREFIX_CANDIDATES}"); sys.exit(1) print(f" Krea prefix: '{krea_prefix}'") # Build block map pairs = build_block_map(zimage_targets, args.remap, KREA_TOTAL_BLOCKS, target_total, args.donor_shift, args.krea_peak, args.zimage_peak) print(f"\nBlock mapping (target -> donor):") for tb, db in pairs[:10]: print(f" z:{tb} -> k:{db}") if len(pairs) > 10: print(f" ... ({len(pairs)} total pairs)") # === Main graft loop === n_modified = 0 for cycle in range(args.cycles): cycle_shift = cycle # simple shift per cycle broadens donor coverage print(f"\nCycle {cycle+1}/{args.cycles} (shift={cycle_shift})") for target_block, donor_block in tqdm(pairs, desc=f"cycle {cycle+1}", unit="blk"): db_actual = (donor_block + cycle_shift) % KREA_TOTAL_BLOCKS # Target prefix (with trailing dot) z_pref = f"{target_prefix_base}.{target_block}." k_pref = f"{krea_prefix}.{db_actual}." # Load Krea tensors for this donor block try: krea_wq = fkrea.get_tensor(k_pref + "attn.wq.weight") krea_wk = fkrea.get_tensor(k_pref + "attn.wk.weight") if args.kv else None krea_wv = fkrea.get_tensor(k_pref + "attn.wv.weight") if args.kv else None krea_wo = fkrea.get_tensor(k_pref + "attn.wo.weight") if args.out_proj else None krea_up = fkrea.get_tensor(k_pref + "mlp.up.weight") if args.mlp else None krea_gate = fkrea.get_tensor(k_pref + "mlp.gate.weight") if (args.mlp and args.include_gate) else None krea_down = fkrea.get_tensor(k_pref + "mlp.down.weight") if (args.mlp and args.mlp_include_w2) else None except Exception as e: print(f" WARNING: could not load Krea block {db_actual}: {e}") continue # QKV graft qkv_out, _ = apply_qkv_graft( zimage_base, z_pref, krea_wq, krea_wk, krea_wv, args.strength, args.kv_strength, args.kv, args.mode, args.orthogonal, source_q_heads, target_q_slots, args.hidden_reduce, dev) for k, v in qkv_out.items(): zimage_base[k] = v n_modified += 1 # out_proj graft if args.out_proj and krea_wo is not None: op_out, _ = apply_out_proj_graft( zimage_base, z_pref, krea_wo, args.out_proj_strength, args.mode, args.orthogonal, source_q_heads, target_q_slots, args.hidden_reduce, dev) for k, v in op_out.items(): zimage_base[k] = v n_modified += 1 # MLP graft if args.mlp and krea_up is not None and target_block in mlp_set: mlp_out, _ = apply_mlp_graft( zimage_base, z_pref, krea_up, krea_gate, krea_down, args.mlp_strength, args.mode, args.orthogonal, args.include_gate, args.mlp_include_w2, args.hidden_reduce, dev) for k, v in mlp_out.items(): zimage_base[k] = v n_modified += 1 print(f"\nFinished {args.cycles} cycle(s). {n_modified} tensor modifications made.") # === Save output === os.makedirs(os.path.dirname(args.output) or ".", exist_ok=True) print(f"\nWriting output: {args.output}") save_file(zimage_base, args.output, metadata=zimage_metadata) print(f"Done. {len(zimage_base)} tensors written.") if __name__ == "__main__": main()