import tempfile
import unittest
from pathlib import Path
from types import SimpleNamespace
from unittest.mock import patch
import torch
from torch import nn
from safetensors.torch import save_file
from diffusion_cli.config import (
GenerationDefaults,
ImageGenerationRequest,
ModelProfile,
ModelSource,
UserConfig,
TOKENIZER_FILES,
buildGenerationConfigFromRequest,
loadUserConfig,
selectModelProfile,
)
from diffusion_cli.krea2_model import (
Krea2Conditioning,
Krea2Config,
Krea2Model,
)
from diffusion_cli.krea2_text_encoder import (
KREA2_MAX_TEXT_TOKENS,
KREA2_PROMPT_PREFIX,
KREA2_PROMPT_SUFFIX,
Krea2TextEncoder,
)
from diffusion_cli.krea2_sampling import sampleKrea2, timesteps
from diffusion_cli.quantization import (
QuantizedLayerSpec,
QuantizationManifest,
ScaledFp8Linear,
)
from diffusion_cli.qwen_image_vae import (
QWEN_IMAGE_LATENT_MEAN,
QwenImageVae,
convertWanVaeToDiffusers,
)
from diffusion_cli.model_inspect import inspectModelSource
class Krea2Test(unittest.TestCase):
def testLongPromptsKeepFixedWindowAndAssistantSuffix(self):
prefix_ids = list(range(34))
suffix_ids = list(range(900, 905))
class FakeTokenizer:
def __call__(
self,
text,
*,
return_tensors,
padding,
truncation,
max_length=None,
add_special_tokens,
):
del return_tensors, truncation, add_special_tokens
value = text[0] if isinstance(text, list) else text
if value == KREA2_PROMPT_PREFIX:
ids = prefix_ids
elif value == KREA2_PROMPT_SUFFIX:
ids = suffix_ids
else:
ids = prefix_ids + list(range(100, 800))
if max_length is not None:
ids = ids[:max_length]
mask = [1] * len(ids)
if padding == "max_length" and max_length is not None:
mask.extend([0] * (max_length - len(ids)))
ids.extend([0] * (max_length - len(ids)))
return {
"input_ids": torch.tensor([ids]),
"attention_mask": torch.tensor([mask]),
}
class FakeModel(nn.Module):
def forward(self, input_ids, **kwargs):
del kwargs
hidden = torch.zeros(
input_ids.shape[0],
input_ids.shape[1],
2560,
)
return SimpleNamespace(hidden_states=[hidden] * 36)
encoder = Krea2TextEncoder(
None,
Path("unused"),
"cpu",
model=FakeModel(),
tokenizer=FakeTokenizer(),
)
conditioning = encoder.encodePrompt("long prompt")
self.assertEqual(
conditioning.hidden_states.shape,
(1, KREA2_MAX_TEXT_TOKENS, 12, 2560),
)
self.assertTrue(conditioning.attention_mask[0, -5:].all())
def testNamedProfilesCoexistWithLegacyConfiguration(self):
with tempfile.TemporaryDirectory() as temp_dir:
root = Path(temp_dir)
config_path = root / "diffusion.toml"
config_path.write_text(
"\n".join(
(
'default_model = "krea2-turbo"',
"",
"[models]",
f'tokenizer = "{root / "tokenizer"}"',
"",
"[model_profiles.krea2-turbo]",
'architecture = "krea2"',
'variant = "turbo"',
f'diffusion_model = "{root / "diffusion.safetensors"}"',
f'text_encoder = "{root / "text.safetensors"}"',
f'vae = "{root / "vae.safetensors"}"',
f'tokenizer = "{root / "tokenizer"}"',
)
),
encoding="utf-8",
)
config = loadUserConfig(config_path)
self.assertIn("legacy", config.model_profiles)
self.assertEqual(selectModelProfile(config).name, "krea2-turbo")
self.assertEqual(config.model_profiles["legacy"].architecture, "z-image")
def testProfileInspectionReportsCompatibility(self):
with tempfile.TemporaryDirectory() as temp_dir:
path = Path(temp_dir) / "diffusion.safetensors"
save_file(
{
"blocks.0.weight": torch.zeros(1),
"txtfusion.projector.weight": torch.zeros(1),
},
path,
)
summary = inspectModelSource(
ModelSource(path, "diffusion_model"),
ModelProfile("krea2-turbo", "krea2", "turbo"),
)
self.assertEqual(summary.architecture_guess, "krea2_diffusion")
self.assertEqual(summary.compatibility, "compatible")
def testTurboProfileDefaults(self):
with tempfile.TemporaryDirectory() as temp_dir:
root = Path(temp_dir)
tokenizer = root / "tokenizer"
tokenizer.mkdir()
for name in TOKENIZER_FILES:
(tokenizer / name).write_text("{}", encoding="utf-8")
profile = ModelProfile(
"krea2-turbo",
"krea2",
"turbo",
diffusion_model=root / "diffusion.safetensors",
text_encoder=root / "text.safetensors",
vae=root / "vae.safetensors",
tokenizer=tokenizer,
)
config = UserConfig(
generation=GenerationDefaults(),
model_profiles={profile.name: profile},
default_model=profile.name,
)
with patch(
"diffusion_cli.config.selectDevice",
return_value=torch.device("cpu"),
), patch(
"diffusion_cli.config.selectDtype",
return_value=torch.float32,
):
generation = buildGenerationConfigFromRequest(
ImageGenerationRequest("a mug"),
config,
)
self.assertEqual(generation.steps, 8)
self.assertEqual(generation.cfg, 0.0)
self.assertEqual(generation.mu, 1.15)
def testScaledFp8LinearMatchesExplicitDequantization(self):
layer = ScaledFp8Linear(2, 1, bias=False, dtype=torch.float32)
stored = torch.tensor([[1.0, 2.0]], dtype=torch.float8_e4m3fn)
layer.weight.data.copy_(stored)
layer.weight_scale.data.fill_(2.0)
value = torch.tensor([[3.0, 4.0]])
expected = value @ (stored.float() * 2.0).transpose(0, 1)
self.assertTrue(torch.allclose(layer(value), expected))
self.assertEqual(layer.weight.dtype, torch.float8_e4m3fn)
def testReducedModelReturnsFiniteVelocity(self):
config = Krea2Config(
features=32,
timestep_width=8,
text_width=16,
heads=4,
kv_heads=2,
blocks=1,
mlp_multiplier=1,
patch_size=2,
latent_channels=4,
text_layers=3,
text_heads=2,
text_kv_heads=2,
)
model = Krea2Model(config, dtype=torch.float32)
conditioning = Krea2Conditioning(
torch.randn(1, 3, 3, 16),
torch.ones(1, 3, dtype=torch.bool),
)
output = model(
torch.randn(1, 4, 4, 4),
torch.tensor([0.5]),
conditioning,
)
self.assertEqual(output.shape, (1, 4, 4, 4))
self.assertTrue(torch.isfinite(output).all())
def testSamplerSkipsUnconditionalAtZeroCfg(self):
calls = []
class FakeModel:
def __call__(self, latent, timestep, conditioning):
calls.append(conditioning)
return torch.ones_like(latent)
conditioning = Krea2Conditioning(
torch.zeros(1, 1, 12, 2560),
torch.ones(1, 1, dtype=torch.bool),
)
output = sampleKrea2(
FakeModel(),
conditioning,
batch_size=1,
height=16,
width=16,
seed=7,
steps=2,
cfg=0.0,
device="cpu",
dtype=torch.float32,
)
self.assertEqual(len(calls), 2)
self.assertEqual(output.shape, (1, 16, 2, 2))
def testTimestepsDescendAndTurboMuIsStable(self):
values = timesteps(16, 8, mu=1.15)
self.assertEqual(len(values), 9)
self.assertEqual(values[0], 1.0)
self.assertEqual(values[-1], 0.0)
self.assertTrue(all(a >= b for a, b in zip(values, values[1:])))
def testVaeDecodeNormalizesTimeAxisAndRange(self):
class Result:
def __init__(self, sample):
self.sample = sample
class FakeVae(nn.Module):
def decode(self, value):
return Result(
torch.zeros(
value.shape[0],
3,
1,
value.shape[3] * 8,
value.shape[4] * 8,
)
)
vae = QwenImageVae(None, "cpu", torch.float32, model=FakeVae())
output = vae.decode(torch.zeros(1, 16, 2, 2))
self.assertEqual(output.shape, (1, 3, 16, 16))
self.assertEqual(float(output.mean()), 0.5)
def testVaeConversionCoversRepresentativeFamilies(self):
state = {
"encoder.conv1.weight": torch.zeros(1),
"decoder.middle.1.proj.weight": torch.zeros(1),
"encoder.downsamples.0.residual.2.weight": torch.zeros(1),
"decoder.upsamples.4.shortcut.weight": torch.zeros(1),
"conv1.weight": torch.zeros(1),
}
converted = convertWanVaeToDiffusers(state)
self.assertIn("encoder.conv_in.weight", converted)
self.assertIn("decoder.mid_block.attentions.0.proj.weight", converted)
self.assertIn("encoder.down_blocks.0.conv1.weight", converted)
self.assertIn(
"decoder.up_blocks.1.resnets.0.conv_shortcut.weight",
converted,
)
self.assertIn("quant_conv.weight", converted)
if __name__ == "__main__":
unittest.main()