#!/usr/bin/env python """Offline tests for the Krea2 transformer port and its native loader handling. - Parity: Krea2Transformer2DModel vs the reference SingleStreamDiT. Builds both from one tiny config, copies the reference state dict into the diffusers port, runs identical inputs, and asserts the forward outputs match. Needs the reference (mmdit.py) at $KREA2_REF_DIR (default /home/ohiom/database/watering-hole). - Zero-init regression: a checkpoint that omits the dormant last.up/last.down residual branch, after materialize_zero_init, produces the same output as the base whose up is zeroed. Port only, so it runs without the reference. - comfy_quant real-file (opt-in): loads an actual ComfyUI int8_tensorwise Krea2 single file through the native loader and verifies SDNQ adoption plus a finite tiny forward. Enabled by setting $KREA2_COMFY_FILE to the .safetensors path; needs network for the base repo config ($KREA2_COMFY_REPO, default CalamitousFelicitousness/Krea-2-Base-Diffusers). No server, no checkpoint (except the opt-in comfy_quant test). """ import importlib.util import os import sys from contextlib import nullcontext from types import SimpleNamespace import torch REF_DIR = os.environ.get("KREA2_REF_DIR", "/home/ohiom/database/watering-hole") CFG = dict( features=128, tdim=32, txtdim=64, heads=4, kvheads=2, multiplier=4, layers=2, patch=2, channels=4, bias=False, theta=1e3, txtlayers=3, txtheads=2, txtkvheads=2, ) def load_reference(): sys.path.insert(0, REF_DIR) import mmdit # The reference pins the cuDNN SDPA backend; neutralize it so both models use the same # default kernel and the test can run on CPU. mmdit.sdpa_kernel = lambda *a, **k: nullcontext() return mmdit def load_port(): path = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "pipelines", "krea2", "transformer_krea2.py")) spec = importlib.util.spec_from_file_location("transformer_krea2", path) mod = importlib.util.module_from_spec(spec) spec.loader.exec_module(mod) return mod def make_inputs(): batch, txtlen, imglen = 2, 5, 9 cdim = CFG["channels"] * CFG["patch"] ** 2 seq = txtlen + imglen gen = torch.Generator().manual_seed(1) img = torch.randn(batch, imglen, cdim, generator=gen) context = torch.randn(batch, txtlen, CFG["txtlayers"], CFG["txtdim"], generator=gen) timestep = torch.rand(batch, generator=gen) pos = torch.randint(0, 16, (batch, seq, 3), generator=gen).float() mask = torch.ones(batch, seq, dtype=torch.bool) mask[0, -2:] = False # exercise the key-padding path return img, context, timestep, pos, mask def run_parity(mmdit, port): torch.manual_seed(0) ref = mmdit.SingleStreamDiT(mmdit.SingleMMDiTConfig(**CFG)).float().eval() mine = port.Krea2Transformer2DModel(**CFG).float().eval() missing, unexpected = mine.load_state_dict(ref.state_dict(), strict=False) assert not missing, f"missing keys when loading reference weights: {missing}" assert not unexpected, f"unexpected keys when loading reference weights: {unexpected}" img, context, timestep, pos, mask = make_inputs() with torch.no_grad(): out_ref = ref(img, context, timestep, pos, mask) out_mine = mine( hidden_states=img, encoder_hidden_states=context, timestep=timestep, position_ids=pos, attention_mask=mask, return_dict=False, )[0] assert out_ref.shape == out_mine.shape, f"shape mismatch: {out_ref.shape} vs {out_mine.shape}" diff = (out_ref - out_mine).abs().max().item() rel = diff / (out_ref.abs().max().item() + 1e-8) print(f"output shape: {tuple(out_mine.shape)}") print(f"max abs diff: {diff:.3e} max rel diff: {rel:.3e}") tol = 1e-4 assert diff < tol, f"PARITY FAILED: max abs diff {diff:.3e} >= {tol}" print("PARITY OK") def bootstrap_repo(): """Make repo modules importable. native_transformer pulls in modules.shared, which needs cmd_args parsed first, so bootstrap it the same way the native-transformer suite does.""" repo = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) if repo not in sys.path: sys.path.insert(0, repo) os.environ.setdefault("SD_INSTALL_QUIET", "1") import modules.cmd_args import installer orig_argv = sys.argv sys.argv = [sys.argv[0]] try: modules.cmd_args.parse_args() finally: sys.argv = orig_argv installer.add_args(modules.cmd_args.parser) modules.cmd_args.parsed, _ = modules.cmd_args.parser.parse_known_args([]) def load_materialize_zero_init(): bootstrap_repo() from pipelines.native_transformer import materialize_zero_init return materialize_zero_init class _FakeTokenizedBatch(SimpleNamespace): def to(self, device): self.input_ids = self.input_ids.to(device) # pylint: disable=attribute-defined-outside-init self.attention_mask = self.attention_mask.to(device) # pylint: disable=attribute-defined-outside-init return self class _FakeTokenizer: def __init__(self, text_len): self.text_len = text_len def __call__(self, texts, truncation=False, padding=None, max_length=None, return_tensors=None): batch = len(texts) seq_len = max_length if padding == "max_length" else self.text_len ids = torch.zeros(batch, seq_len, dtype=torch.int64) mask = torch.zeros(batch, seq_len, dtype=torch.bool) for i, text in enumerate(texts): length = min(len(text), seq_len) mask[i, :length] = True return _FakeTokenizedBatch(input_ids=ids, attention_mask=mask) class _FakeTextEncoder: def __init__(self, hidden_dim, layers): self.hidden_dim = hidden_dim self.layers = layers def __call__(self, input_ids=None, attention_mask=None, output_hidden_states=False): batch, seq_len = input_ids.shape hidden_states = [torch.randn(batch, seq_len, self.hidden_dim) for _ in range(self.layers)] return SimpleNamespace(hidden_states=hidden_states) def run_dense_prompt_compaction_test(): """Validate SD_KREA2_DENSE prompt compaction and mask preservation.""" os.environ["SD_KREA2_DENSE"] = "1" from pipelines.krea2.pipeline_krea2 import Krea2Pipeline class FakeTransformer: def __init__(self): self.dtype = torch.float32 self.config = SimpleNamespace(patch=2, channels=4) fake_transformer = FakeTransformer() fake_text_encoder = _FakeTextEncoder(hidden_dim=32, layers=36) fake_tokenizer = _FakeTokenizer(text_len=20) fake_vae = SimpleNamespace(config=SimpleNamespace(latents_mean=0.0, latents_std=1.0), decode=lambda x: SimpleNamespace(sample=x)) fake_scheduler = SimpleNamespace() pipe = Krea2Pipeline( transformer=fake_transformer, text_encoder=fake_text_encoder, tokenizer=fake_tokenizer, vae=fake_vae, scheduler=fake_scheduler, ) prompts = ["short", "longer prompt"] _hidden, mask = pipe.encode_prompt(prompts, device=torch.device("cpu")) assert mask.shape[1] < pipe.MAX_LENGTH assert mask.any(dim=0).all(), "Compacted prompt mask must contain no fully padded columns" del os.environ["SD_KREA2_DENSE"] def run_attention_mask_none_forward_test(port): """Verify the transformer accepts attention_mask=None without dense segment mask creation.""" model = port.Krea2Transformer2DModel(**CFG).eval() batch, txtlen, imglen = 1, 5, 9 cdim = CFG["channels"] * CFG["patch"] ** 2 img = torch.randn(batch, imglen, cdim) context = torch.randn(batch, txtlen, CFG["txtlayers"], CFG["txtdim"]) timestep = torch.rand(batch) pos = torch.randint(0, 16, (batch, txtlen + imglen, 3)).float() with torch.no_grad(): out = model( hidden_states=img, encoder_hidden_states=context, timestep=timestep, position_ids=pos, attention_mask=None, return_dict=False, )[0] assert out.shape == (batch, imglen, CFG["channels"] * CFG["patch"] ** 2) def run_zero_init_regression(port): """A finetune predating the last.up/last.down branch omits both keys. Zero-filling them must reproduce the base's dormant behavior: identical output to a model whose up is zeroed (up(down(x)) == 0), while a genuinely missing weight stays a hard mismatch.""" materialize_zero_init = load_materialize_zero_init() branch_keys = ("last.up.weight", "last.down.weight") torch.manual_seed(0) full_sd = port.Krea2Transformer2DModel(**CFG).state_dict() # Ground truth: the branch dormant exactly as the base ships it (up zeroed, down arbitrary). dormant = port.Krea2Transformer2DModel(**CFG).float().eval() dormant_sd = dict(full_sd) dormant_sd["last.up.weight"] = torch.zeros_like(dormant_sd["last.up.weight"]) dormant.load_state_dict(dormant_sd) # Under test: a checkpoint that omits both branch keys; the loader zero-fills them. A real # missing weight (blocks.0.attn.wq.weight) is dropped too, to confirm it is NOT zero-filled. filled = port.Krea2Transformer2DModel(**CFG).float().eval() hard_key = "blocks.0.attn.wq.weight" partial_sd = {k: v for k, v in full_sd.items() if k not in branch_keys and k != hard_key} missing, unexpected = filled.load_state_dict(partial_sd, strict=False) assert not unexpected, f"unexpected keys: {unexpected}" assert set(missing) == set(branch_keys) | {hard_key}, f"unexpected missing set: {missing}" remaining = materialize_zero_init(filled, missing, ("last.down.", "last.up.")) assert remaining == [hard_key], f"expected only {hard_key} to remain hard-missing, got {remaining}" for k in branch_keys: w = dict(filled.named_parameters())[k] assert w.abs().sum().item() == 0.0, f"{k} was not zero-filled" # Restore the genuinely-missing weight so the forward is well defined, then compare. filled.load_state_dict({hard_key: full_sd[hard_key]}, strict=False) img, context, timestep, pos, mask = make_inputs() def forward(model): with torch.no_grad(): return model( hidden_states=img, encoder_hidden_states=context, timestep=timestep, position_ids=pos, attention_mask=mask, return_dict=False, )[0] diff = (forward(dormant) - forward(filled)).abs().max().item() print(f"zero-init regression: max abs diff vs dormant base: {diff:.3e}") assert diff == 0.0, f"ZERO-INIT REGRESSION FAILED: filled output differs by {diff:.3e}" print("ZERO-INIT OK") def run_comfy_quant_real_file(): """Opt-in end-to-end check against a real ComfyUI int8_tensorwise Krea2 file: the native loader must adopt every marked linear as an SDNQ int8 layer and produce a finite output on a tiny forward. $KREA2_COMFY_LAYERS overrides the expected layer count (default 224).""" path = os.environ.get("KREA2_COMFY_FILE") if not path: print("COMFY REAL-FILE SKIPPED (set KREA2_COMFY_FILE to enable)") return assert os.path.exists(path), f"KREA2_COMFY_FILE not found: {path}" expected_layers = int(os.environ.get("KREA2_COMFY_LAYERS", "224")) repo_id = os.environ.get("KREA2_COMFY_REPO", "CalamitousFelicitousness/Krea-2-Base-Diffusers") bootstrap_repo() from pipelines import native_transformer as nt from pipelines.krea2 import KREA2_SPEC transformer, siblings = nt.load(local_file=path, repo_id=repo_id, spec=KREA2_SPEC, diffusers_cfg={}) assert not siblings sdnq_layers = [m for m in transformer.modules() if m.__class__.__name__ == "SDNQLinear"] storage_dtypes = {m.weight.dtype for m in sdnq_layers} print(f"comfy_quant real file: {len(sdnq_layers)} SDNQ layers, storage {storage_dtypes}") assert len(sdnq_layers) == expected_layers, f"expected {expected_layers} SDNQ layers, got {len(sdnq_layers)}" assert storage_dtypes <= {torch.int8, torch.float8_e4m3fn, torch.uint8}, f"unexpected storage dtypes: {storage_dtypes}" assert len(storage_dtypes) == 1, "adopted weights must share one storage dtype" assert transformer.blocks[0].attn.wq.__class__.__name__ == "SDNQLinear" assert getattr(transformer, "quantization_config", None) is not None cfg = transformer.config param = next(p for p in transformer.parameters() if p.is_floating_point()) device, dtype = param.device, param.dtype batch, txtlen, imglen = 1, 3, 4 seq = txtlen + imglen gen = torch.Generator().manual_seed(1) img = torch.randn(batch, imglen, cfg.channels * cfg.patch ** 2, generator=gen).to(device=device, dtype=dtype) context = torch.randn(batch, txtlen, cfg.txtlayers, cfg.txtdim, generator=gen).to(device=device, dtype=dtype) timestep = torch.rand(batch, generator=gen).to(device=device, dtype=dtype) pos = torch.randint(0, 16, (batch, seq, 3), generator=gen).float().to(device=device) mask = torch.ones(batch, seq, dtype=torch.bool, device=device) with torch.no_grad(): out = transformer( hidden_states=img, encoder_hidden_states=context, timestep=timestep, position_ids=pos, attention_mask=mask, return_dict=False, )[0] assert torch.isfinite(out).all(), "forward produced non-finite values" print("COMFY REAL-FILE OK") def main(): port = load_port() # Parity runs first in pristine torch state; the regression imports modules afterwards. mmdit = load_reference() run_parity(mmdit, port) run_zero_init_regression(port) run_comfy_quant_real_file() if __name__ == "__main__": main()