mirror of
https://github.com/vladmandic/automatic
synced 2026-09-12 16:08:43 +02:00
aa8cd9980e
Signed-off-by: Vladimir Mandic <mandic00@live.com>
314 lines
13 KiB
Python
314 lines
13 KiB
Python
#!/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()
|