#!/usr/bin/env python3 """Ming-Image speed + fidelity harness: one model load, N prompts. Reuses infer.py's own loader and generation path unchanged. The only addition is a wrapper around model.diffusion_loss.sample that records the conditioning tensors the DiT receives (encoder_hidden_states / directvlm_hidden_states) and times the sampling stage (DiT steps + VAE decode) separately from the MLLM stage. usage: ming_bench.py --prompts a.json b.json --out DIR [--repeat-first N] -- are passed to infer.parse_args() as-is (e.g. --model, --resolution, --steps, --seed, --device, --device-map none, --attn-implementation eager, --int8-mllm). --repeat-first N re-runs the first prompt N more times at the same seed: the images measure the platform's run-to-run noise floor and the timings are warm timings. """ import argparse import json import sys import time from pathlib import Path def main(): ap = argparse.ArgumentParser() ap.add_argument("--prompts", nargs="+", required=True) ap.add_argument("--out", required=True) ap.add_argument("--repeat-first", type=int, default=0) own, rest = ap.parse_known_args() if rest and rest[0] == "--": rest = rest[1:] sys.argv = [sys.argv[0], "--prompt", own.prompts[0]] + rest import torch from safetensors.torch import save_file import infer args = infer.parse_args() model_directory = infer.resolve_model_directory( args.model, revision=args.revision, cache_dir=args.cache_dir, local_files_only=args.local_files_only, ) profile = infer.load_checkpoint_capabilities(model_directory) resolution = infer.resolve_task_resolution(args.task, args.resolution) sampling = profile.resolve_sampling_parameters(steps=args.steps, cfg=args.cfg) dtype = infer._dtype(args.dtype) out = Path(own.out) out.mkdir(parents=True, exist_ok=True) def sync(): if torch.cuda.is_available(): torch.cuda.synchronize() sync() t0 = time.perf_counter() model, processor = infer.load_model_and_processor(model_directory, args) sync() load_s = time.perf_counter() - t0 print(f"LOAD_S {load_s:.1f}", flush=True) captured = {} original_sample = model.diffusion_loss.sample def recording_sample(*a, **kw): for key in ("encoder_hidden_states", "directvlm_hidden_states"): value = kw.get(key) if isinstance(value, (list, tuple)): value = torch.stack(list(value), dim=0) if isinstance(value, torch.Tensor): captured[key] = value.detach().float().cpu().contiguous() sync() ts = time.perf_counter() result = original_sample(*a, **kw) sync() captured["_sample_s"] = time.perf_counter() - ts return result model.diffusion_loss.sample = recording_sample runs = [(p, 0) for p in own.prompts] + [(own.prompts[0], i + 1) for i in range(own.repeat_first)] results = [] for prompt_path, rep in runs: stem = Path(prompt_path).stem + (f"_rep{rep}" if rep else "") prompt = infer._load_prompt(prompt_path) captured.clear() if torch.cuda.is_available(): torch.cuda.reset_peak_memory_stats() sync() t1 = time.perf_counter() images = infer.run_generation( model, processor, profile, task=args.task, prompt=prompt, input_image=None, resolution=resolution, sampling=sampling, seed=args.seed, num_layers=args.num_layers, dtype=dtype, ) sync() total_s = time.perf_counter() - t1 if len(images) != 1: raise RuntimeError(f"{stem}: expected 1 image, got {len(images)}") image_path = out / f"{stem}.png" images[0].save(image_path) cond = {k: v for k, v in captured.items() if not k.startswith("_")} if "encoder_hidden_states" not in cond: raise RuntimeError(f"{stem}: conditioning was not captured") save_file(cond, str(out / f"{stem}.cond.safetensors")) sample_s = captured["_sample_s"] row = { "load_s": round(load_s, 1), "prompt": str(prompt_path), "stem": stem, "seed": args.seed, "resolution": resolution, "steps": sampling.steps, "cfg": sampling.cfg, "mode": images[0].mode, "size": list(images[0].size), "total_s": round(total_s, 2), "sample_s": round(sample_s, 2), "mllm_s": round(total_s - sample_s, 2), "peak_alloc_gib": round(torch.cuda.max_memory_allocated() / 2**30, 2) if torch.cuda.is_available() else None, "cond_shapes": {k: list(v.shape) for k, v in cond.items()}, } results.append(row) print("RUN " + json.dumps(row), flush=True) with open(out / "runs.jsonl", "a") as fh: # accumulates across one-prompt-per-process runs fh.write(json.dumps(row) + "\n") manifest = {"load_s": round(load_s, 1), "args": {k: str(v) for k, v in vars(args).items()}, "runs": results} (out / "manifest.json").write_text(json.dumps(manifest, indent=2)) print("BENCH_DONE", out, flush=True) if __name__ == "__main__": main()