WaveCut's picture
Finalize quality and inference optimized OrbitQuant release
f84e201 verified
Raw
History Blame Contribute Delete
9.78 kB
#!/usr/bin/env python3
"""Generate with Krea 2 Turbo while staging Qwen and DiT to minimize VRAM."""
from __future__ import annotations
import argparse
import gc
import json
import os
from pathlib import Path
from typing import Any
import torch
os.environ.setdefault("ORBITQUANT_STRICT_PACKED", "1")
import orbitquant # noqa: F401 - register OrbitQuant with Hugging Face loaders.
from orbitquant.layers import OrbitQuantLinear
SELECTED_LAYERS = (2, 5, 8, 11, 14, 17, 20, 23, 26, 29, 32, 35)
PROMPT_PREFIX = (
"<|im_start|>system\nDescribe the image by detailing the color, shape, size, "
"texture, quantity, text, spatial relationships of the objects and "
"background:<|im_end|>\n<|im_start|>user\n"
)
PROMPT_SUFFIX = "<|im_end|>\n<|im_start|>assistant\n"
PROMPT_PREFIX_TOKENS = 34
PROMPT_SUFFIX_TOKENS = 5
def install_strict_flash_attention() -> None:
"""Use Flash SDPA for Krea's all-valid DiT attention and fail on fallback."""
from diffusers.models import attention_dispatch
from torch.nn.attention import SDPBackend, sdpa_kernel
def attention(
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_mask: torch.Tensor | None = None,
dropout_p: float = 0.0,
is_causal: bool = False,
scale: float | None = None,
enable_gqa: bool = False,
return_lse: bool = False,
_parallel_config: Any | None = None,
) -> torch.Tensor:
if return_lse:
raise ValueError("strict Flash attention does not support return_lse=True")
if _parallel_config is not None:
raise ValueError("strict Flash attention does not support context parallelism")
query, key, value = (
tensor.permute(0, 2, 1, 3) for tensor in (query, key, value)
)
with sdpa_kernel(SDPBackend.FLASH_ATTENTION):
output = torch.nn.functional.scaled_dot_product_attention(
query,
key,
value,
attn_mask=None,
dropout_p=dropout_p,
is_causal=is_causal,
scale=scale,
enable_gqa=enable_gqa,
)
return output.permute(0, 2, 1, 3)
backend = attention_dispatch.AttentionBackendName.NATIVE
attention_dispatch._AttentionBackendRegistry._backends[backend] = attention
attention_dispatch._AttentionBackendRegistry._supported_arg_names[backend] = {
"query",
"key",
"value",
"attn_mask",
"dropout_p",
"is_causal",
"scale",
"enable_gqa",
"return_lse",
"_parallel_config",
}
attention_dispatch._AttentionBackendRegistry.set_active_backend(backend)
def orbit_inventory(model: torch.nn.Module) -> dict[str, Any]:
modules = [module for module in model.modules() if isinstance(module, OrbitQuantLinear)]
return {
"orbitquant_linear_count": len(modules),
"executed_orbitquant_linear_count": sum(
module.last_effective_runtime_mode is not None for module in modules
),
"effective_runtime_modes": sorted(
{
module.last_effective_runtime_mode
for module in modules
if module.last_effective_runtime_mode is not None
}
),
"shared_activation_cache_hit_count": sum(
bool(getattr(module, "last_activation_cache_hit", False))
for module in modules
),
"full_dequantized_cache_count": sum(
module._dequantized_weight_cache is not None for module in modules
),
}
def compact_prompt_embeddings(
embeddings: torch.Tensor, mask: torch.Tensor
) -> tuple[torch.Tensor, torch.Tensor]:
"""Remove padded prompt lanes before using unmasked Flash attention."""
if embeddings.ndim != 4 or mask.ndim != 2 or embeddings.shape[:2] != mask.shape:
raise ValueError("prompt embeddings and mask have incompatible shapes")
if embeddings.shape[0] != 1:
raise ValueError("lossless prompt compaction currently requires batch size 1")
valid = mask[0].bool()
if not bool(valid.any()):
raise ValueError("prompt contains no valid tokens")
compacted = embeddings[:, valid]
compacted_mask = torch.ones(
(1, compacted.shape[1]), dtype=torch.bool, device=mask.device
)
return compacted, compacted_mask
@torch.no_grad()
def encode_prompt(
model: str, prompt: str, max_sequence_length: int, revision: str | None
):
from transformers import Qwen2Tokenizer, Qwen3VLModel
tokenizer = Qwen2Tokenizer.from_pretrained(
model, subfolder="tokenizer", revision=revision
)
encoder = Qwen3VLModel.from_pretrained(
model, subfolder="text_encoder", revision=revision, dtype=torch.bfloat16
).eval().requires_grad_(False).to("cuda")
text = tokenizer(
[PROMPT_PREFIX + prompt],
truncation=True,
padding="max_length",
max_length=max_sequence_length + PROMPT_PREFIX_TOKENS - PROMPT_SUFFIX_TOKENS,
return_tensors="pt",
).to("cuda")
suffix = tokenizer([PROMPT_SUFFIX], return_tensors="pt").to("cuda")
input_ids = torch.cat([text.input_ids, suffix.input_ids], dim=1)
attention_mask = torch.cat([text.attention_mask, suffix.attention_mask], dim=1).bool()
position_ids = (attention_mask.long().cumsum(dim=-1) - 1).clamp(min=0)
outputs = encoder(
input_ids=input_ids,
attention_mask=attention_mask,
position_ids=position_ids.unsqueeze(0).expand(3, -1, -1),
output_hidden_states=True,
)
embeddings = torch.stack(
[outputs.hidden_states[index] for index in SELECTED_LAYERS], dim=2
)[:, PROMPT_PREFIX_TOKENS:].cpu()
mask = attention_mask[:, PROMPT_PREFIX_TOKENS:].cpu()
inventory = orbit_inventory(encoder)
embeddings, mask = compact_prompt_embeddings(embeddings, mask)
del outputs, encoder, tokenizer, text, suffix, input_ids, attention_mask, position_ids
gc.collect()
torch.cuda.empty_cache()
return embeddings, mask, inventory
def main() -> int:
parser = argparse.ArgumentParser()
parser.add_argument("--model", default="WaveCut/Krea-2-Turbo-OrbitQuant-W4A4")
parser.add_argument("--revision")
parser.add_argument("--prompt", required=True)
parser.add_argument("--output", type=Path, default=Path("krea2-orbitquant.png"))
parser.add_argument("--width", type=int, default=2048)
parser.add_argument("--height", type=int, default=2048)
parser.add_argument("--steps", type=int, default=8)
parser.add_argument("--seed", type=int, default=0)
parser.add_argument("--max-sequence-length", type=int, default=512)
parser.add_argument("--vae-tile-size", type=int, default=1024)
parser.add_argument("--vae-tile-stride", type=int, default=896)
args = parser.parse_args()
install_strict_flash_attention()
embeddings, mask, qwen_inventory = encode_prompt(
args.model, args.prompt, args.max_sequence_length, args.revision
)
from diffusers import (
AutoencoderKLQwenImage,
FlowMatchEulerDiscreteScheduler,
Krea2Pipeline,
Krea2Transformer2DModel,
)
scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained(
args.model, subfolder="scheduler", revision=args.revision
)
vae = AutoencoderKLQwenImage.from_pretrained(
args.model,
subfolder="vae",
revision=args.revision,
torch_dtype=torch.bfloat16,
).eval().requires_grad_(False).to("cuda")
vae.enable_tiling(
tile_sample_min_height=args.vae_tile_size,
tile_sample_min_width=args.vae_tile_size,
tile_sample_stride_height=args.vae_tile_stride,
tile_sample_stride_width=args.vae_tile_stride,
)
transformer = Krea2Transformer2DModel.from_pretrained(
args.model,
subfolder="transformer",
revision=args.revision,
torch_dtype=torch.bfloat16,
).eval().requires_grad_(False).to("cuda")
pipe = Krea2Pipeline(
scheduler=scheduler,
vae=vae,
text_encoder=None,
tokenizer=None,
transformer=transformer,
text_encoder_select_layers=SELECTED_LAYERS,
is_distilled=True,
patch_size=2,
)
image = pipe(
prompt_embeds=embeddings.to("cuda"),
prompt_embeds_mask=mask.to("cuda"),
width=args.width,
height=args.height,
num_inference_steps=args.steps,
guidance_scale=0.0,
max_sequence_length=args.max_sequence_length,
generator=torch.Generator(device="cuda").manual_seed(args.seed),
).images[0]
dit_inventory = orbit_inventory(transformer)
for name, inventory in (("Qwen", qwen_inventory), ("DiT", dit_inventory)):
if inventory["effective_runtime_modes"] != ["native_packed_matmul"]:
raise RuntimeError(f"{name} did not use packed OrbitQuant: {inventory}")
if inventory["full_dequantized_cache_count"]:
raise RuntimeError(f"{name} retained full dequantized weight caches")
if not inventory["shared_activation_cache_hit_count"]:
raise RuntimeError(
f"{name} did not reuse activation quantization across adjacent projections"
)
args.output.parent.mkdir(parents=True, exist_ok=True)
image.save(args.output)
print(
json.dumps(
{
"output": str(args.output),
"qwen": qwen_inventory,
"dit": dit_inventory,
"torch_peak_mb": torch.cuda.max_memory_allocated() / (1024**2),
},
indent=2,
)
)
return 0
if __name__ == "__main__":
raise SystemExit(main())