Mage-Flow-NVFP4-AJH / generate.py
ajh-code's picture
Add files using upload-large-folder tool
54152e6 verified
Raw
History Blame Contribute Delete
14.5 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 quant_text_encoder import load_quantized_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_quantized_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,
) -> 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:
module.to("cpu")
torch.cuda.synchronize()
torch.cuda.empty_cache()
metrics["seconds"] = time.perf_counter() - started
return value, metrics
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,
) -> 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] = {}
def encode_text() -> tuple[Any, ...]:
verdict = model.txt_enc.screen_text(prompt)
if verdict.violates:
return (verdict,)
text_flat, vec_all, text_lens = _encode_texts_packed(
model,
[prompt, negative_prompt or " "],
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)
encoded, stage_metrics["text_encode"] = stage_to_gpu(
model.txt_enc,
encode_text,
torch,
)
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,
)
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,
)
screening = {
"violates": bool(verdict.violates),
"categories": list(verdict.categories or []),
"reason": str(verdict.reason),
}
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=20)
parser.add_argument("--cfg", type=float, default=5.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())