File size: 30,803 Bytes
c6ac3cc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
"""radiance_infer.py β€” Minimal standalone inference script for the Radiance model.

──────────────────────────────────────────────────────────────────────────────
QUICK START
──────────────────────────────────────────────────────────────────────────────

    python radiance_infer.py \\
        --ckpt  latest_x0_full_20M_dataset_run.safetensors \\
        --t5    models/t5xxl \\
        --prompt "a cinematic shot of a red fox in autumn forest"

All defaults are tuned for quality:
    cfg=3.5  steps=50  mu=1.0  resolution=1024Γ—1024  batch=1

──────────────────────────────────────────────────────────────────────────────
MODEL ARCHITECTURE OVERVIEW  (for backend implementors)
──────────────────────────────────────────────────────────────────────────────

Radiance is a **pixel-space flow-matching image generator** based on MM-DiT
(multimodal Diffusion Transformer).  Everything is in the pixel domain β€” there
is no VAE or latent space.

  β”Œβ”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”
  β”‚  Input: noisy RGB image [B, 3, H, W]  +  T5 text embedding β”‚
  β”‚                                                            β”‚
  β”‚  1. Patchify via Conv2d (patch_size=16, zero-init)         β”‚
  β”‚     β†’ image token sequence  [B, N, hidden]                 β”‚
  β”‚                                                            β”‚
  β”‚  2. Text projection (Linear 4096 β†’ hidden)                 β”‚
  β”‚     β†’ text token sequence   [B, L, hidden]                 β”‚
  β”‚                                                            β”‚
  β”‚  3. Approximator (distilled AdaLN modulation network)      β”‚
  β”‚     Given (timestep, guidance=0), generates ALL shift/     β”‚
  β”‚     scale/gate vectors for every transformer block in      β”‚
  β”‚     one shot.  Runs under torch.no_grad() β€” intentional.   β”‚
  β”‚                                                            β”‚
  β”‚  4. depth Γ— DoubleStreamBlock (MM-DiT)                     β”‚
  β”‚     Parallel cross-attention over image + text streams     β”‚
  β”‚     with 3-axis RoPE (time, height, width).                β”‚
  β”‚                                                            β”‚
  β”‚  5. depth_single_blocks Γ— SingleStreamBlock (DiT)          β”‚
  β”‚     Merged image+text stream.                              β”‚
  β”‚                                                            β”‚
  β”‚  6. NeRF decoder head                                      β”‚
  β”‚     Per-patch hypernetwork (NerfGLUBlock Γ— nerf_depth)     β”‚
  β”‚     conditioned on the transformer output.  Reconstructs   β”‚
  β”‚     x0 at full pixel resolution via a 3Γ—3 conv fold.       β”‚
  β”‚                                                            β”‚
  β”‚  Output: predicted x0 [B, 3, H, W] β†’ converted to          β”‚
  β”‚          v-prediction  v = (x_noisy βˆ’ x0) / (t + Ξ΅)        β”‚
  β””β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”˜

FLOW MATCHING & V-PREDICTION
──────────────────────────────────────────────────────────────────────────────

Training uses flow-matching with a straight-line interpolation between clean
data x1 and Gaussian noise x0:

    x_noisy(t) = t * x_noise + (1-t) * x_clean     t ∈ [0, 1]

The model predicts the velocity v such that moving along v brings x_noisy
towards x_clean.  At inference, Euler steps integrate the ODE from t=1 to t=0:

    x_{t-dt} = x_t + (t_next - t_curr) * v(x_t, t_curr)

CLASSIFIER-FREE GUIDANCE (CFG)
──────────────────────────────────────────────────────────────────────────────

At each Euler step, two forward passes are performed:

    v_pos = model(x, t, positive_text_embedding)
    v_neg = model(x, t, negative_text_embedding)   # typically empty string

    v_cfg = v_neg + cfg_scale * (v_pos - v_neg)

cfg_scale=1.0 disables guidance.  Values 3–7 are typical.

SCHEDULE MU  (schedule_mu parameter)
──────────────────────────────────────────────────────────────────────────────

`schedule_mu` controls how Euler timesteps are distributed over [0, 1]:

    mu = None   β†’ auto-shift based on image sequence length (recommended)
    mu = 0.0    β†’ uniform linear spacing
    mu = float  β†’ shifted via CDF inversion of a parabolic density:

        p(t) ∝ -7.7·(t - 0.5)² + 2   shifted by logit-space transform:
        t_shifted = t / (t + (1-t)Β·exp(-ΞΌ))

    mu > 0 β†’ more steps concentrated at low noise (high t) β†’ more detail
    mu < 0 β†’ more steps concentrated at high noise (low t) β†’ better structure
    mu = 1.0 (default) is a good balance for 1024Γ—1024.

TEXT ENCODER
──────────────────────────────────────────────────────────────────────────────

Radiance uses T5-XXL (11B parameters, 4096-dim hidden) as a frozen text
encoder.  Only the encoder side is loaded (no decoder).  Embeddings are
produced at bfloat16 with padding to `max_seq_len` (default 512 tokens).

A blank-string embedding is used as the unconditional negative for CFG.
Passing your own `--neg_prompt` allows soft negative conditioning.

──────────────────────────────────────────────────────────────────────────────
DEPENDENCIES  (beyond PyTorch)
──────────────────────────────────────────────────────────────────────────────

    pip install safetensors transformers einops tqdm pillow

This script also requires the `src/` directory from this repo to be importable.
Add it to your PYTHONPATH or run from the repo root:

    PYTHONPATH=/path/to/x0-pred python radiance_infer.py ...

──────────────────────────────────────────────────────────────────────────────
"""
from __future__ import annotations

import argparse
import json
import os
import sys
from pathlib import Path

import torch
from safetensors.torch import load_file as load_safetensors

# ── src/ must be on the path ──────────────────────────────────────────────────
# When running from the repo root this is automatic; otherwise add it explicitly.
_REPO_ROOT = Path(__file__).parent
if str(_REPO_ROOT) not in sys.path:
    sys.path.insert(0, str(_REPO_ROOT))

from src.models.radiance import Radiance, RadianceParams  # noqa: E402


# ──────────────────────────────────────────────────────────────────────────────
# Optional: transformers (T5 text encoder)
# ──────────────────────────────────────────────────────────────────────────────
try:
    from transformers import AutoTokenizer, T5EncoderModel
    _TRANSFORMERS_AVAILABLE = True
except ImportError:
    _TRANSFORMERS_AVAILABLE = False


# ──────────────────────────────────────────────────────────────────────────────
# Helpers
# ──────────────────────────────────────────────────────────────────────────────

def _strip_compiled_keys(sd: dict) -> dict:
    """Remove the ``_orig_mod.`` prefix that torch.compile adds to state-dict keys.

    When a model is saved after torch.compile(), every key in the state dict
    gains an ``_orig_mod.`` prefix.  This strips it so the weights can be
    loaded into an uncompiled model instance.
    """
    prefix = "_orig_mod."
    return {k.replace(prefix, "") if prefix in k else k: v for k, v in sd.items()}


def _encode_text(
    text_encoder: "T5EncoderModel",
    tokenizer: "AutoTokenizer",
    texts: list[str],
    max_seq_len: int,
    device: str,
) -> tuple[torch.Tensor, torch.Tensor]:
    """Encode a list of strings with a frozen T5EncoderModel.

    Pads / truncates to `max_seq_len`.  Everything runs in bfloat16 under
    torch.no_grad() β€” the encoder is never updated during inference.

    Args:
        text_encoder: HuggingFace T5EncoderModel (encoder-only, bfloat16).
        tokenizer:    Matching AutoTokenizer.
        texts:        List of B prompt strings.
        max_seq_len:  Token budget (512 by default, matches training).
        device:       Target device string, e.g. "cuda" or "cuda:0".

    Returns:
        embeddings: float tensor [B, max_seq_len, 4096]  β€” last hidden states.
        mask:       bool tensor  [B, max_seq_len]        β€” 1 for real tokens.
    """
    inputs = tokenizer(
        texts,
        padding="max_length",
        max_length=max_seq_len,
        truncation=True,
        return_tensors="pt",
    ).to(device)

    with torch.no_grad(), torch.autocast(device_type="cuda", dtype=torch.bfloat16):
        outputs = text_encoder(
            input_ids=inputs.input_ids,
            attention_mask=inputs.attention_mask,
        )

    embeddings = outputs.last_hidden_state   # [B, L, 4096]
    mask       = inputs.attention_mask.bool()
    return embeddings, mask


# ──────────────────────────────────────────────────────────────────────────────
# Model loading
# ──────────────────────────────────────────────────────────────────────────────

def load_model(
    ckpt_path: str,
    config_path: str,
    device: str,
) -> Radiance:
    """Load and return a Radiance model ready for inference.

    The config JSON must contain a ``model_config`` dict that matches the
    fields in ``RadianceParams``.  The simplest way is to pass the
    ``config.json`` that lives alongside the checkpoint directory.

    Weights are loaded from a safetensors file (.safetensors or .sft).
    ``_orig_mod.`` prefixes from torch.compile are stripped automatically.

    The model is moved to `device`, cast to bfloat16, and set to eval mode.
    Gradient checkpointing is disabled (not needed at inference).

    Args:
        ckpt_path:   Path to the .safetensors checkpoint file.
        config_path: Path to the JSON config file.
        device:      Target device, e.g. "cuda" or "cuda:1".

    Returns:
        Radiance model in eval mode on the requested device.
    """
    print(f"[radiance_infer] Loading config from: {config_path}")
    with open(config_path) as f:
        cfg = json.load(f)

    model_cfg = cfg["model_config"]
    # Disable grad checkpointing at inference β€” only needed during training
    model_cfg = {**model_cfg, "grad_checkpointing": False}

    params = RadianceParams(**model_cfg)
    model  = Radiance(params)

    print(f"[radiance_infer] Loading weights from: {ckpt_path}")
    state_dict = load_safetensors(ckpt_path, device="cpu")
    state_dict = _strip_compiled_keys(state_dict)

    missing, unexpected = model.load_state_dict(state_dict, strict=True)
    if missing:
        print(f"[radiance_infer] WARNING: {len(missing)} missing keys: {missing[:5]} ...")
    if unexpected:
        print(f"[radiance_infer] WARNING: {len(unexpected)} unexpected keys: {unexpected[:5]} ...")

    model = model.to(device=device, dtype=torch.bfloat16)
    model.eval()
    print(f"[radiance_infer] Model ready on {device}  (bfloat16).")
    return model


def load_text_encoder(
    t5_path: str,
    device: str,
    max_seq_len: int,
) -> tuple["T5EncoderModel", "AutoTokenizer"]:
    """Load the frozen T5-XXL encoder and its tokenizer.

    Only the encoder stack is loaded β€” no decoder, no language model head.
    Parameters are frozen (requires_grad=False) and the model is set to eval.

    T5-XXL produces 4096-dim hidden states, which is what Radiance expects
    via its ``context_in_dim`` parameter.  If you swap to a different text
    encoder (e.g. Qwen3-2560), adjust ``context_in_dim`` in the model config.

    Args:
        t5_path:     Directory containing the HuggingFace T5 encoder weights.
        device:      Target device string.
        max_seq_len: Token budget β€” used only to print a note; the tokenizer
                     will be invoked with this value at encode time.

    Returns:
        (T5EncoderModel, AutoTokenizer) both ready for use.
    """
    if not _TRANSFORMERS_AVAILABLE:
        raise RuntimeError(
            "transformers is required for text encoding. "
            "Install it with:  pip install transformers"
        )

    print(f"[radiance_infer] Loading T5 tokenizer from: {t5_path}")
    tokenizer = AutoTokenizer.from_pretrained(t5_path)

    print(f"[radiance_infer] Loading T5 encoder from: {t5_path}  (bfloat16, device={device})")
    encoder = T5EncoderModel.from_pretrained(t5_path, torch_dtype=torch.bfloat16)
    encoder = encoder.to(device).eval()
    for p in encoder.parameters():
        p.requires_grad_(False)

    print(f"[radiance_infer] T5 encoder ready.  Max token budget: {max_seq_len}.")
    return encoder, tokenizer


# ──────────────────────────────────────────────────────────────────────────────
# Inference
# ──────────────────────────────────────────────────────────────────────────────

def run_inference(
    model: Radiance,
    encoder: "T5EncoderModel",
    tokenizer: "AutoTokenizer",
    prompts: list[str],
    neg_prompt: str,
    cfg_scale: float,
    num_steps: int,
    schedule_mu: float | None,
    width: int,
    height: int,
    device: str,
    max_seq_len: int,
    seed: int,
    output_dir: str,
) -> list[str]:
    """Run CFG Euler sampling and save images to disk.

    Each prompt produces one image (batch size derived from len(prompts)).
    Images are saved as ``{output_dir}/{i:04d}_{prompt_slug}.png``.

    Args:
        model:       Radiance model (eval, bfloat16, on device).
        encoder:     Frozen T5EncoderModel.
        tokenizer:   Matching AutoTokenizer.
        prompts:     List of positive prompt strings β€” one image per prompt.
        neg_prompt:  Single negative conditioning string, broadcast to all images.
        cfg_scale:   CFG guidance scale.  1.0 = no guidance, 3–7 = typical.
        num_steps:   Number of Euler integration steps.  28–50 is typical.
        schedule_mu: Timestep schedule shift (see module docstring).
                       None  β†’ auto (recommended).
                       0.0   β†’ linear uniform.
                       float β†’ parabolic CDF shift.
        width:       Output image width in pixels (must be divisible by 16).
        height:      Output image height in pixels (must be divisible by 16).
        device:      PyTorch device string.
        max_seq_len: T5 token budget.
        seed:        RNG seed for reproducible noise.
        output_dir:  Directory to write output PNG files.

    Returns:
        List of saved file paths.
    """
    # ── Validate resolution ────────────────────────────────────────────────────
    if width % 16 != 0 or height % 16 != 0:
        raise ValueError(
            f"Resolution ({width}Γ—{height}) must be divisible by 16 "
            f"(the model patch size is 16)."
        )

    os.makedirs(output_dir, exist_ok=True)
    B = len(prompts)

    # ── Encode text ────────────────────────────────────────────────────────────
    print(f"[radiance_infer] Encoding {B} prompt(s)...")
    pos_embeds, pos_mask = _encode_text(encoder, tokenizer, prompts,       max_seq_len, device)
    neg_embeds, neg_mask = _encode_text(encoder, tokenizer, [neg_prompt]*B, max_seq_len, device)
    # Shapes: [B, max_seq_len, 4096] and [B, max_seq_len]

    # ── Sample initial noise ───────────────────────────────────────────────────
    # Flow-matching starts at t=1 (pure Gaussian noise) and integrates to t=0
    # (clean image).  Using a fixed seed makes results reproducible.
    generator = torch.Generator(device=device).manual_seed(seed)
    noise = torch.randn(
        B, 3, height, width,
        dtype=torch.bfloat16,
        device=device,
        generator=generator,
    )

    # ── Run Euler CFG sampling ─────────────────────────────────────────────────
    print(
        f"[radiance_infer] Sampling  "
        f"cfg={cfg_scale}  steps={num_steps}  mu={schedule_mu}  "
        f"{width}Γ—{height}  seed={seed}"
    )
    with torch.autocast("cuda", torch.bfloat16):
        images, _ = model.euler_cfg(
            x           = noise,
            cfg_scale   = cfg_scale,
            num_steps   = num_steps,
            txt         = pos_embeds,
            txt_mask    = pos_mask,
            neg_txt     = neg_embeds,
            neg_txt_mask= neg_mask,
            schedule_mu = schedule_mu,   # None β†’ auto-mu from seq length
        )
    # images: [B, 3, H, W]  float bfloat16  in [-1, 1]

    # ── Save images ────────────────────────────────────────────────────────────
    # Rescale from [-1, 1] β†’ [0, 255] uint8, then save as PNG via PIL.
    try:
        from PIL import Image
        import numpy as np
        _USE_PIL = True
    except ImportError:
        _USE_PIL = False

    saved_paths = []
    images_f32 = images.float().clamp(-1.0, 1.0)   # ensure no out-of-range values

    for i, (img_t, prompt) in enumerate(zip(images_f32, prompts)):
        # img_t: [3, H, W]  in [-1, 1]
        img_01  = (img_t + 1.0) / 2.0               # [0, 1]
        img_u8  = (img_01 * 255.0).byte().cpu()      # [3, H, W] uint8

        # Build a filename from the prompt (truncated, spaces β†’ underscores)
        slug = prompt[:60].replace(" ", "_").replace("/", "-")
        fname = f"{i:04d}_{slug}.png"
        fpath = os.path.join(output_dir, fname)

        if _USE_PIL:
            # PIL expects HWC layout
            arr = img_u8.permute(1, 2, 0).numpy()
            Image.fromarray(arr, mode="RGB").save(fpath)
        else:
            # Fallback: raw bytes via torch (requires torchvision for PNG, but
            # this at least saves something if PIL is absent)
            try:
                from torchvision.io import write_png
                write_png(img_u8, fpath)
            except ImportError:
                raise RuntimeError(
                    "Neither pillow nor torchvision is available for saving images. "
                    "Install one:  pip install pillow"
                )

        saved_paths.append(fpath)
        print(f"[radiance_infer]   Saved: {fpath}")

    return saved_paths


# ──────────────────────────────────────────────────────────────────────────────
# CLI entry point
# ──────────────────────────────────────────────────────────────────────────────

def _build_parser() -> argparse.ArgumentParser:
    p = argparse.ArgumentParser(
        prog="radiance_infer.py",
        description=(
            "Minimal inference script for the Radiance pixel-space flow-matching model.\n"
            "Runs CFG Euler sampling from T=1 (noise) to T=0 (image) and saves PNG files."
        ),
        formatter_class=argparse.ArgumentDefaultsHelpFormatter,
    )

    # ── Required ──────────────────────────────────────────────────────────────
    p.add_argument(
        "--ckpt", required=True,
        help="Path to the .safetensors model checkpoint.",
    )
    p.add_argument(
        "--prompt", required=True, nargs="+",
        help=(
            "One or more prompt strings.  Each prompt generates one image.  "
            "Use quotes for multi-word prompts: --prompt 'a red fox' 'a blue sky'."
        ),
    )

    # ── Text encoder ──────────────────────────────────────────────────────────
    p.add_argument(
        "--t5", default=None,
        help=(
            "Path to the T5-XXL encoder directory (HuggingFace format).  "
            "Falls back to the 't5_path' field in the config JSON."
        ),
    )
    p.add_argument(
        "--neg_prompt", default="",
        help=(
            "Negative conditioning text, broadcast to all images.  "
            "An empty string (default) is the standard unconditional baseline."
        ),
    )
    p.add_argument(
        "--max_seq_len", type=int, default=512,
        help="T5 token budget.  Prompts longer than this are truncated.",
    )

    # ── Config ────────────────────────────────────────────────────────────────
    p.add_argument(
        "--config", default=None,
        help=(
            "Path to the model config JSON.  Auto-detected in this order:\n"
            "  1. config.json in the checkpoint's parent directory\n"
            "  2. config_radiance.json in the current working directory\n"
            "  3. config.json in the current working directory"
        ),
    )

    # ── Sampling hyperparameters ──────────────────────────────────────────────
    p.add_argument(
        "--cfg", type=float, default=3.5,
        help=(
            "CFG guidance scale.  "
            "1.0 = no guidance (unconditional).  "
            "3–7 = typical range.  Higher = stronger prompt adherence."
        ),
    )
    p.add_argument(
        "--steps", type=int, default=50,
        help=(
            "Number of Euler integration steps.  "
            "More steps = slower but potentially cleaner output.  "
            "28 is a fast setting; 50 is the quality default."
        ),
    )
    p.add_argument(
        "--mu", type=float, default=1.0,
        help=(
            "Schedule shift strength (schedule_mu).  "
            "Controls timestep density distribution over [0, 1]:\n"
            "  > 0 β†’ more steps at low noise (fine detail focus)\n"
            "  < 0 β†’ more steps at high noise (global structure focus)\n"
            "  0.0 β†’ uniform linear spacing\n"
            "Set to 'auto' to use sequence-length-based auto-mu "
            "(pass --mu 0 and --auto_mu instead if needed)."
        ),
    )
    p.add_argument(
        "--auto_mu", action="store_true",
        help=(
            "Use automatic schedule_mu derived from image sequence length "
            "(overrides --mu).  This is the Flux/Chroma default behaviour."
        ),
    )

    # ── Resolution & batch ────────────────────────────────────────────────────
    p.add_argument("--width",  type=int, default=1024,
                   help="Output image width in pixels (must be divisible by 16).")
    p.add_argument("--height", type=int, default=1024,
                   help="Output image height in pixels (must be divisible by 16).")
    p.add_argument(
        "--batch", type=int, default=1,
        help=(
            "Number of images to generate per prompt.  "
            "If >1, the same prompt is repeated `batch` times with different seeds."
        ),
    )

    # ── Output & reproducibility ──────────────────────────────────────────────
    p.add_argument("--output", default="output",
                   help="Directory to write output PNG files.")
    p.add_argument("--seed",   type=int, default=42,
                   help="RNG seed for reproducible noise initialisation.")
    p.add_argument("--device", default="cuda",
                   help="PyTorch device.  Multi-GPU not supported here; use 'cuda:N'.")

    return p


def _resolve_config(args_config: str | None, ckpt_path: str) -> str:
    """Auto-detect the config JSON path if not explicitly provided."""
    if args_config is not None:
        if not os.path.isfile(args_config):
            raise FileNotFoundError(f"Config file not found: {args_config}")
        return args_config

    candidates = [
        # 1. config.json in the checkpoint's parent directory
        os.path.join(os.path.dirname(ckpt_path), "config.json"),
        # 2. config_radiance.json in cwd
        "config_radiance.json",
        # 3. config.json in cwd
        "config.json",
    ]
    for c in candidates:
        if os.path.isfile(c):
            print(f"[radiance_infer] Auto-detected config: {c}")
            return c

    raise FileNotFoundError(
        "Could not find a config JSON.  Pass --config explicitly.\n"
        f"Tried: {candidates}"
    )


def _resolve_t5(args_t5: str | None, config_path: str) -> str:
    """Resolve the T5 path from CLI arg or config JSON."""
    if args_t5 is not None:
        return args_t5
    with open(config_path) as f:
        cfg = json.load(f)
    t5 = cfg.get("t5_path") or cfg.get("tokenizer_path")
    if not t5:
        raise ValueError(
            "T5 path not found in config.  Pass --t5 /path/to/t5xxl explicitly."
        )
    return t5


def main() -> None:
    parser = _build_parser()
    args   = parser.parse_args()

    # ── Resolve paths ──────────────────────────────────────────────────────────
    config_path = _resolve_config(args.config, args.ckpt)
    t5_path     = _resolve_t5(args.t5, config_path)

    # ── Expand batch > 1: repeat each prompt `batch` times ───────────────────
    prompts = []
    for prompt in args.prompt:
        prompts.extend([prompt] * args.batch)
    # Stagger seeds across repeats so they don't all look identical
    # (seed is used per-batch; varying per image requires a separate loop,
    # but for batch > 1 we just bump the seed for each extra copy)

    # ── Build schedule_mu ─────────────────────────────────────────────────────
    # None β†’ auto (get_schedule), float β†’ parabolic-CDF shifted, 0.0 β†’ linear
    schedule_mu: float | None = None if args.auto_mu else args.mu

    # ── Load model & text encoder ─────────────────────────────────────────────
    model   = load_model(args.ckpt, config_path, args.device)
    encoder, tokenizer = load_text_encoder(t5_path, args.device, args.max_seq_len)

    # ── Inference ─────────────────────────────────────────────────────────────
    saved = run_inference(
        model        = model,
        encoder      = encoder,
        tokenizer    = tokenizer,
        prompts      = prompts,
        neg_prompt   = args.neg_prompt,
        cfg_scale    = args.cfg,
        num_steps    = args.steps,
        schedule_mu  = schedule_mu,
        width        = args.width,
        height       = args.height,
        device       = args.device,
        max_seq_len  = args.max_seq_len,
        seed         = args.seed,
        output_dir   = args.output,
    )

    print(f"\n[radiance_infer] Done. {len(saved)} image(s) saved to '{args.output}/'.")


if __name__ == "__main__":
    main()