ajh-code's picture
Use single transformer branch at Turbo CFG 1
1553346 verified
Raw
History Blame Contribute Delete
16.1 kB
#!/usr/bin/env python3
"""Portable single-image Mage-Flow NVFP4 inference entry point."""
from __future__ import annotations
import argparse
from datetime import datetime, timezone
import gc
import json
import os
from pathlib import Path
import platform
import sys
import time
from typing import Any, Callable
RELEASE_ROOT = Path(__file__).resolve().parent
RUNTIME_ROOT = RELEASE_ROOT / "runtime"
VENDOR_ROOT = RELEASE_ROOT / "vendor"
DEFAULT_MODEL = str(RELEASE_ROOT)
for import_root in (RUNTIME_ROOT, VENDOR_ROOT):
if str(import_root) not in sys.path:
sys.path.insert(0, str(import_root))
def utc_now() -> str:
return (
datetime.now(timezone.utc)
.replace(microsecond=0)
.isoformat()
.replace("+00:00", "Z")
)
def resolve_model(value: str) -> Path:
candidate = Path(value).expanduser()
if candidate.is_dir():
return candidate.resolve()
from huggingface_hub import snapshot_download
return Path(
snapshot_download(
repo_id=value,
allow_patterns=[
"model_index.json",
"transformer/config.json",
"transformer/*.safetensors",
"transformer/*.json",
"text_encoder/*",
"vae/config.json",
"vae/*.safetensors",
"scheduler/*",
],
)
).resolve()
def _repo_subpath(repo_dir: Path, relative: str) -> str:
path = (repo_dir / relative).resolve()
if not path.is_relative_to(repo_dir):
raise ValueError(f"model path escapes its snapshot: {relative}")
return str(path)
def _structure_from_config(transformer_config: dict[str, Any]) -> dict[str, Any]:
metadata_keys = {
"_class_name",
"txt_max_length",
"max_sequence_length",
"param_dtype",
"packing",
"schedule_mode",
"static_shift",
"use_time_shift",
"rope_type",
"apply_text_rotary_emb",
"mlp_ratio",
"depth_single_blocks",
"theta",
"qkv_bias",
"guidance_embed",
"vec_in_dim",
"vec_type",
"time_type",
"double_block_type",
"quantization_config",
}
return {
key: value
for key, value in transformer_config.items()
if key not in metadata_keys
}
def load_pipeline(
*,
model: str,
torch: Any,
) -> tuple[Any, dict[str, Any]]:
import torch.nn as nn
from diffusers import FlowMatchEulerDiscreteScheduler
from mage_flow.models.mage_flow import MageFlowModel, ModelConfig
from mage_flow.models.modules._attn_backend import set_attn_backend
from mage_flow.pipeline import MageFlowPipeline
from text_encoder_variants import load_scaled_fp8_text_encoder
from standard_transformer import load_standard_native_transformer
repo_dir = resolve_model(model)
model_index_path = repo_dir / "model_index.json"
transformer_config_path = repo_dir / "transformer" / "config.json"
model_index = json.loads(model_index_path.read_text(encoding="utf-8"))
transformer_config = json.loads(
transformer_config_path.read_text(encoding="utf-8")
)
structure = _structure_from_config(transformer_config)
config = ModelConfig(
vae_path=_repo_subpath(repo_dir, model_index["_vae_source"]),
txt_enc_path=_repo_subpath(repo_dir, model_index["_text_encoder_path"]),
model_structure=structure,
txt_max_length=transformer_config.get("txt_max_length", 2048),
packing=transformer_config.get("packing", True),
static_shift=transformer_config.get("static_shift", 6.0),
)
transformer, access_report = load_standard_native_transformer(
repo_dir,
torch.device("cuda:0"),
)
model = MageFlowModel.__new__(MageFlowModel)
nn.Module.__init__(model)
model.config = config
set_attn_backend(getattr(config, "attn_type", "flash2"))
model.patch_text_encoder_forward()
model.vae = model.load_vae()
model.transformer = transformer
text_encoder_dir = Path(
_repo_subpath(repo_dir, model_index["_text_encoder_path"])
)
model.txt_enc, text_report = load_scaled_fp8_text_encoder(
text_encoder_dir=text_encoder_dir,
artifact_path=text_encoder_dir / "model.safetensors",
tokenizer_max_length=config.txt_max_length,
dit_structure=structure,
use_packed_text_infer=config.packing,
)
model.vae.requires_grad_(False).to(torch.bfloat16)
model.txt_enc.requires_grad_(False)
model.eval()
model.scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained(
_repo_subpath(repo_dir, "scheduler")
)
return (
MageFlowPipeline(model, device="cuda:0"),
{
"resolved_model": str(repo_dir),
"transformer_access": access_report,
"text_encoder": text_report,
},
)
def stage_to_gpu(
module: Any,
operation: Callable[[], Any],
torch: Any,
*,
offload_after: bool = True,
) -> tuple[Any, dict[str, Any]]:
torch.cuda.synchronize()
torch.cuda.reset_peak_memory_stats(0)
started = time.perf_counter()
try:
module.to("cuda:0")
value = operation()
torch.cuda.synchronize()
metrics = {
"peak_allocated_bytes": int(torch.cuda.max_memory_allocated(0)),
"peak_reserved_bytes": int(torch.cuda.max_memory_reserved(0)),
}
finally:
if offload_after:
module.to("cpu")
torch.cuda.synchronize()
torch.cuda.empty_cache()
metrics["seconds"] = time.perf_counter() - started
return value, metrics
def place_pipeline_resident(
pipe: Any,
torch: Any,
) -> dict[str, Any]:
"""Place the text encoder, transformer, and VAE on one GPU for reuse."""
torch.cuda.synchronize()
started = time.perf_counter()
for module in (
pipe.model.txt_enc,
pipe.model.transformer,
pipe.model.vae,
):
module.to("cuda:0")
torch.cuda.synchronize()
free_bytes, total_bytes = torch.cuda.mem_get_info(0)
return {
"seconds": time.perf_counter() - started,
"allocated_bytes": int(torch.cuda.memory_allocated(0)),
"reserved_bytes": int(torch.cuda.memory_reserved(0)),
"free_bytes": int(free_bytes),
"total_bytes": int(total_bytes),
}
def generate_staged(
*,
pipe: Any,
prompt: str,
negative_prompt: str,
height: int,
width: int,
steps: int,
cfg: float,
seed: int,
static_shift: float,
torch: Any,
keep_resident: bool = False,
) -> tuple[Any, dict[str, Any], dict[str, Any]]:
from einops import rearrange
from mage_flow.models.modules.mage_latent import encode_noise, resolve_gs_key
from mage_flow.pipeline import (
_build_pack_ctx,
_decode_one,
_encode_texts_packed,
_get_scheduler,
_lens_to_cu,
_make_divisible_by_16,
_slice_packed,
_template_info,
_velocity,
)
model = pipe.model
device = torch.device("cuda:0")
template_info = _template_info("mage-flow")
template = template_info.get("template", "{}")
drop_index = int(template_info.get("start_idx", 0))
stage_metrics: dict[str, Any] = {}
use_negative = cfg > 1.0 and bool(negative_prompt.strip())
def encode_text() -> tuple[Any, ...]:
verdict = model.txt_enc.screen_text(prompt)
if verdict.violates:
return (verdict,)
if use_negative:
text_flat, vec_all, text_lens = _encode_texts_packed(
model,
[prompt, negative_prompt],
template,
drop_index,
device,
)
positive = _slice_packed(
text_flat, vec_all, text_lens, 0, 1, device
)
negative = _slice_packed(
text_flat, vec_all, text_lens, 1, 1, device
)
return (verdict, *positive, *negative)
text_flat, vec_all, text_lens = _encode_texts_packed(
model,
[prompt],
template,
drop_index,
device,
)
positive = _slice_packed(
text_flat, vec_all, text_lens, 0, 1, device
)
return (verdict, *positive, None, None, None, None)
encoded, stage_metrics["text_encode"] = stage_to_gpu(
model.txt_enc,
encode_text,
torch,
offload_after=not keep_resident,
)
verdict = encoded[0]
if verdict.violates:
raise RuntimeError(
"prompt was refused by Mage content screening: "
f"{verdict.categories} {verdict.reason}"
)
(
_verdict,
txt,
txt_cu,
txt_mask,
vec,
neg_txt,
neg_cu,
neg_mask,
neg_vec,
) = encoded
height = _make_divisible_by_16(height)
width = _make_divisible_by_16(width)
noise = encode_noise(
(
model.vae.latent_channels,
(height + 15) // 16,
(width + 15) // 16,
),
key=resolve_gs_key(None),
seed=seed,
device=device,
dtype=torch.bfloat16,
)
_, _, grid_h, grid_w = noise.shape
image_latent = rearrange(noise, "b c h w -> b (h w) c")
image_ids = torch.zeros(grid_h, grid_w, 3, device=device)
image_ids[..., 1] += torch.arange(grid_h, device=device)[:, None]
image_ids[..., 2] += torch.arange(grid_w, device=device)[None, :]
image_ids = rearrange(image_ids, "h w c -> 1 (h w) c")
image_lens = [grid_h * grid_w]
image_cu = _lens_to_cu(image_lens, device)
context = _build_pack_ctx(
image_ids,
image_cu,
[[(1, grid_h, grid_w)]],
image_lens,
txt,
txt_cu,
txt_mask,
vec,
neg_txt,
neg_cu,
neg_mask,
neg_vec,
cfg,
False,
True,
device,
)
def denoise() -> Any:
nonlocal image_latent
scheduler = _get_scheduler(
model,
steps,
device,
static_shift,
)
for step_index, timestep in enumerate(scheduler.timesteps):
prediction = _velocity(
model.transformer,
image_latent,
context,
scheduler.sigmas[step_index].item(),
)
image_latent = scheduler.step(
prediction,
timestep,
image_latent,
return_dict=False,
)[0]
if int((~torch.isfinite(image_latent)).sum().item()) != 0:
raise RuntimeError("denoising produced non-finite values")
return image_latent.to("cpu")
latent_cpu, stage_metrics["denoise"] = stage_to_gpu(
model.transformer,
denoise,
torch,
offload_after=not keep_resident,
)
def decode() -> Any:
return _decode_one(
model,
latent_cpu.to(device),
height,
width,
device,
)
image, stage_metrics["decode"] = stage_to_gpu(
model.vae,
decode,
torch,
offload_after=not keep_resident,
)
screening = {
"violates": bool(verdict.violates),
"categories": list(verdict.categories or []),
"reason": str(verdict.reason),
"cfg_branch_count": 2 if use_negative else 1,
"negative_prompt_used": use_negative,
}
return image, stage_metrics, screening
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--prompt", required=True)
parser.add_argument("--negative-prompt", default="")
parser.add_argument("--output", type=Path, default=Path("mage_nvfp4.png"))
parser.add_argument(
"--model",
default=DEFAULT_MODEL,
help=(
"local standard-layout model directory or Hugging Face repo id "
"(default: this downloaded repository)"
),
)
parser.add_argument("--height", type=int, default=1024)
parser.add_argument("--width", type=int, default=1024)
parser.add_argument("--steps", type=int, default=4)
parser.add_argument("--cfg", type=float, default=1.0)
parser.add_argument("--seed", type=int, default=1)
parser.add_argument("--static-shift", type=float, default=6.0)
return parser.parse_args()
def main() -> int:
args = parse_args()
output_path = args.output.expanduser().resolve()
report_path = output_path.with_suffix(output_path.suffix + ".json")
if output_path.exists() or report_path.exists():
raise SystemExit(
f"refusing to overwrite existing output/report: {output_path}"
)
output_path.parent.mkdir(parents=True, exist_ok=True)
import torch
from packed_nvfp4_linear import close_all_contexts
from torch_ops_native import (
close_native_contexts,
initialize_native_sm120_op,
)
if not torch.cuda.is_available() or torch.cuda.device_count() != 1:
raise SystemExit(
"exactly one visible CUDA GPU is required; set CUDA_VISIBLE_DEVICES"
)
torch.cuda.set_device(0)
properties = torch.cuda.get_device_properties(0)
if (properties.major, properties.minor) != (12, 0):
raise SystemExit(
f"native NVFP4 requires SM120; found {properties.major}.{properties.minor}"
)
if not initialize_native_sm120_op(allow_python_schema_fallback=False):
raise SystemExit("the packaged native SM120 torch op did not load")
os.environ.setdefault("CUBLAS_WORKSPACE_CONFIG", ":4096:8")
torch.manual_seed(args.seed)
torch.cuda.manual_seed_all(args.seed)
torch.backends.cudnn.benchmark = False
torch.backends.cudnn.deterministic = True
torch.backends.cuda.matmul.allow_tf32 = False
torch.use_deterministic_algorithms(True)
started = time.perf_counter()
pipe = None
try:
pipe, load_report = load_pipeline(
model=args.model,
torch=torch,
)
image, stages, screening = generate_staged(
pipe=pipe,
prompt=args.prompt,
negative_prompt=args.negative_prompt,
height=args.height,
width=args.width,
steps=args.steps,
cfg=args.cfg,
seed=args.seed,
static_shift=args.static_shift,
torch=torch,
)
image.save(output_path)
report = {
"schema_version": "mage-flow-nvfp4-portable-generation-v1",
"status": "success",
"completed_at_utc": utc_now(),
"output": str(output_path),
"prompt": args.prompt,
"negative_prompt": args.negative_prompt,
"height": args.height,
"width": args.width,
"steps": args.steps,
"cfg": args.cfg,
"seed": args.seed,
"elapsed_seconds": time.perf_counter() - started,
"stages": stages,
"screening": screening,
"load": load_report,
"environment": {
"python": platform.python_version(),
"torch": torch.__version__,
"torch_cuda": torch.version.cuda,
"gpu": properties.name,
"compute_capability": (
f"{properties.major}.{properties.minor}"
),
},
}
report_path.write_text(
json.dumps(report, indent=2, sort_keys=True) + "\n",
encoding="utf-8",
)
print(f"saved {output_path}")
print(f"saved {report_path}")
return 0
finally:
close_native_contexts()
close_all_contexts()
pipe = None
gc.collect()
if __name__ == "__main__":
raise SystemExit(main())