Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions .gitignore
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
__pycache__/
.mypy_cache/
.vscode
.cursor
*.engine
Expand Down
4 changes: 1 addition & 3 deletions examples/visual_gen/visual_gen_ltx2.py
Original file line number Diff line number Diff line change
Expand Up @@ -383,8 +383,6 @@ def main():

start_time = time.time()

inputs = {"prompt": args.prompt}

extra_params = {
"guidance_rescale": args.guidance_rescale,
"stg_scale": args.stg_scale,
Expand All @@ -411,7 +409,7 @@ def main():
extra_params=extra_params,
)

output = visual_gen.generate(inputs=inputs, params=params)
output = visual_gen.generate(inputs=args.prompt, params=params)

end_time = time.time()
logger.info(f"Generation completed in {end_time - start_time:.2f}s")
Expand Down
2 changes: 1 addition & 1 deletion examples/visual_gen/visual_gen_wan_i2v.py
Original file line number Diff line number Diff line change
Expand Up @@ -336,7 +336,7 @@ def main():
extra_params["boundary_ratio"] = args.boundary_ratio

output = visual_gen.generate(
inputs={"prompt": args.prompt},
inputs=args.prompt,
params=VisualGenParams(
height=args.height,
width=args.width,
Expand Down
2 changes: 1 addition & 1 deletion examples/visual_gen/visual_gen_wan_t2v.py
Original file line number Diff line number Diff line change
Expand Up @@ -340,7 +340,7 @@ def main():
extra_params["boundary_ratio"] = args.boundary_ratio

output = visual_gen.generate(
inputs={"prompt": args.prompt},
inputs=args.prompt,
params=VisualGenParams(
height=args.height,
width=args.width,
Expand Down
78 changes: 36 additions & 42 deletions tensorrt_llm/_torch/visual_gen/executor.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@
import threading
import traceback
from dataclasses import dataclass
from typing import List, Optional, Union
from typing import TYPE_CHECKING, List, Optional

import torch
import torch.distributed as dist
Expand All @@ -15,42 +15,24 @@
from tensorrt_llm.executor.ipc import ZeroMqQueue
from tensorrt_llm.logger import logger

if TYPE_CHECKING:
from tensorrt_llm.visual_gen.params import VisualGenParams


@dataclass
class DiffusionRequest:
"""Request for diffusion inference.

Universal parameters are top-level fields with ``None`` meaning
"use model default" (resolved by the executor before calling
``pipeline.infer()``). Model-specific parameters live in
``extra_params`` and are passed through to the pipeline.
Generation parameters live in the optional ``params`` object
(a :class:`~tensorrt_llm.visual_gen.params.VisualGenParams` instance).
When ``params`` is ``None`` (the default), the executor creates a
``VisualGenParams()`` and fills it with pipeline-specific defaults
before calling ``pipeline.infer()``.
"""

request_id: int
prompt: List[str]
negative_prompt: Optional[str] = None

# Core — None means "use model default" (resolved by executor)
height: Optional[int] = None
width: Optional[int] = None
num_inference_steps: Optional[int] = None
guidance_scale: Optional[float] = None
max_sequence_length: Optional[int] = None
seed: int = 42

# Video
num_frames: Optional[int] = None
frame_rate: Optional[float] = None

# Image
num_images_per_prompt: int = 1

# Conditioning inputs
image: Optional[Union[str, bytes, List[Union[str, bytes]]]] = None
image_cond_strength: Optional[float] = None

# Model-specific overflow (from VisualGenParams.extra_params)
extra_params: Optional[dict] = None
params: Optional["VisualGenParams"] = None


@dataclass
Expand Down Expand Up @@ -216,29 +198,40 @@ def serve_forever(self):
self.process_request(req)

def _merge_defaults(self, req: DiffusionRequest):
"""Fill ``None`` fields in *req* with pipeline-specific defaults.
"""Fill ``None`` fields in *req.params* with pipeline-specific defaults.

Merges both universal defaults (from ``default_generation_params``)
and extra_param defaults (from ``extra_param_specs``).
"""
if req.params is None:
from tensorrt_llm.visual_gen.params import VisualGenParams

kwargs = dict(self.pipeline.default_generation_params)
specs = self.pipeline.extra_param_specs
if specs:
kwargs["extra_params"] = {key: spec.default for key, spec in specs.items()}
req.params = VisualGenParams(**kwargs)
return

params = req.params
# Universal field defaults
for field_name, default_value in self.pipeline.default_generation_params.items():
if hasattr(req, field_name) and getattr(req, field_name) is None:
setattr(req, field_name, default_value)
if hasattr(params, field_name) and getattr(params, field_name) is None:
setattr(params, field_name, default_value)

# Extra param defaults — fill all declared keys so infer() can use direct access
specs = self.pipeline.extra_param_specs
if specs:
if req.extra_params is None:
req.extra_params = {}
if params.extra_params is None:
params.extra_params = {}
for key, spec in specs.items():
if key not in req.extra_params:
req.extra_params[key] = spec.default
if key not in params.extra_params:
params.extra_params[key] = spec.default

self._validate_request(req)

def _validate_request(self, req: DiffusionRequest):
"""Validate *req* against the loaded pipeline's declared parameters.
"""Validate *req.params* against the loaded pipeline's declared parameters.

Raises ``VisualGenParamsError`` on:
- Unknown ``extra_params`` keys
Expand All @@ -251,14 +244,15 @@ def _validate_request(self, req: DiffusionRequest):
# (executor → visual_gen.visual_gen → _torch.visual_gen → executor)
from tensorrt_llm.visual_gen.visual_gen import VisualGenParamsError

params = req.params
errors: list[str] = []
pipeline_name = self.pipeline.__class__.__name__
declared_defaults = self.pipeline.default_generation_params
specs = self.pipeline.extra_param_specs

# --- unknown extra_params keys ---
if req.extra_params:
unknown = set(req.extra_params.keys()) - set(specs.keys())
if params.extra_params:
unknown = set(params.extra_params.keys()) - set(specs.keys())
if unknown:
errors.append(
f"Unknown extra_params {sorted(unknown)} for {pipeline_name}. "
Expand All @@ -271,7 +265,7 @@ def _validate_request(self, req: DiffusionRequest):
# Conditioning inputs (image, negative_prompt, mask) are excluded —
# they are validated at runtime by the pipeline's infer().
for field_name in _GENERATION_CONFIG_FIELDS:
value = getattr(req, field_name, None)
value = getattr(params, field_name, None)
if value is not None and field_name not in declared_defaults:
errors.append(
f"Parameter '{field_name}' is set but {pipeline_name} does "
Expand All @@ -280,8 +274,8 @@ def _validate_request(self, req: DiffusionRequest):
)

# --- extra_params type and range checks ---
if req.extra_params:
for key, value in req.extra_params.items():
if params.extra_params:
for key, value in params.extra_params.items():
if key not in specs:
continue # already reported as unknown above
spec = specs[key]
Expand Down Expand Up @@ -315,7 +309,7 @@ def process_request(self, req: DiffusionRequest):
try:
self._merge_defaults(req)
cache_key = self.pipeline.warmup_cache_key(
req.height, req.width, num_frames=req.num_frames
req.params.height, req.params.width, num_frames=req.params.num_frames
)
if self.pipeline._warmed_up_shapes and cache_key not in self.pipeline._warmed_up_shapes:
logger.warning(
Expand Down
12 changes: 6 additions & 6 deletions tensorrt_llm/_torch/visual_gen/models/flux/pipeline_flux.py
Original file line number Diff line number Diff line change
Expand Up @@ -246,12 +246,12 @@ def infer(self, req):
"""Run inference from DiffusionRequest."""
return self.forward(
prompt=req.prompt,
height=req.height,
width=req.width,
num_inference_steps=req.num_inference_steps,
guidance_scale=req.guidance_scale,
seed=req.seed,
max_sequence_length=req.max_sequence_length,
height=req.params.height,
width=req.params.width,
num_inference_steps=req.params.num_inference_steps,
guidance_scale=req.params.guidance_scale,
seed=req.params.seed,
max_sequence_length=req.params.max_sequence_length,
)

@torch.inference_mode()
Expand Down
12 changes: 6 additions & 6 deletions tensorrt_llm/_torch/visual_gen/models/flux/pipeline_flux2.py
Original file line number Diff line number Diff line change
Expand Up @@ -339,12 +339,12 @@ def infer(self, req):
"""Run inference from DiffusionRequest."""
return self.forward(
prompt=req.prompt,
height=req.height,
width=req.width,
num_inference_steps=req.num_inference_steps,
guidance_scale=req.guidance_scale,
seed=req.seed,
max_sequence_length=req.max_sequence_length,
height=req.params.height,
width=req.params.width,
num_inference_steps=req.params.num_inference_steps,
guidance_scale=req.params.guidance_scale,
seed=req.params.seed,
max_sequence_length=req.params.max_sequence_length,
)

@torch.inference_mode()
Expand Down
24 changes: 12 additions & 12 deletions tensorrt_llm/_torch/visual_gen/models/ltx2/pipeline_ltx2.py
Original file line number Diff line number Diff line change
Expand Up @@ -1170,22 +1170,22 @@ def extra_param_specs(self):

def infer(self, req):
"""Run inference with request parameters."""
extra = req.extra_params or {}
extra = req.params.extra_params or {}
return self.forward(
prompt=req.prompt,
negative_prompt=req.negative_prompt,
height=req.height,
width=req.width,
num_frames=req.num_frames,
frame_rate=req.frame_rate,
num_inference_steps=req.num_inference_steps,
guidance_scale=req.guidance_scale,
seed=req.seed,
negative_prompt=req.params.negative_prompt,
height=req.params.height,
width=req.params.width,
num_frames=req.params.num_frames,
frame_rate=req.params.frame_rate,
num_inference_steps=req.params.num_inference_steps,
guidance_scale=req.params.guidance_scale,
seed=req.params.seed,
output_type=extra["output_type"],
guidance_rescale=extra["guidance_rescale"],
max_sequence_length=req.max_sequence_length,
image=req.image,
image_cond_strength=req.image_cond_strength,
max_sequence_length=req.params.max_sequence_length,
image=req.params.image,
image_cond_strength=req.params.image_cond_strength,
stg_scale=extra["stg_scale"],
stg_blocks=extra["stg_blocks"],
modality_scale=extra["modality_scale"],
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -565,22 +565,22 @@ def load_standard_components(
# ------------------------------------------------------------------

def infer(self, req):
extra = req.extra_params or {}
extra = req.params.extra_params or {}
return self.forward(
prompt=req.prompt,
negative_prompt=req.negative_prompt,
height=req.height,
width=req.width,
num_frames=req.num_frames,
frame_rate=req.frame_rate,
num_inference_steps=req.num_inference_steps,
guidance_scale=req.guidance_scale,
seed=req.seed,
negative_prompt=req.params.negative_prompt,
height=req.params.height,
width=req.params.width,
num_frames=req.params.num_frames,
frame_rate=req.params.frame_rate,
num_inference_steps=req.params.num_inference_steps,
guidance_scale=req.params.guidance_scale,
seed=req.params.seed,
output_type=extra["output_type"],
guidance_rescale=extra["guidance_rescale"],
max_sequence_length=req.max_sequence_length,
image=req.image,
image_cond_strength=req.image_cond_strength,
max_sequence_length=req.params.max_sequence_length,
image=req.params.image,
image_cond_strength=req.params.image_cond_strength,
stg_scale=extra["stg_scale"],
stg_blocks=extra["stg_blocks"],
modality_scale=extra["modality_scale"],
Expand Down
18 changes: 9 additions & 9 deletions tensorrt_llm/_torch/visual_gen/models/wan/pipeline_wan.py
Original file line number Diff line number Diff line change
Expand Up @@ -325,19 +325,19 @@ def extra_param_specs(self):

def infer(self, req):
"""Run inference with request parameters."""
extra = req.extra_params or {}
extra = req.params.extra_params or {}
return self.forward(
prompt=req.prompt,
negative_prompt=req.negative_prompt,
height=req.height,
width=req.width,
num_frames=req.num_frames,
num_inference_steps=req.num_inference_steps,
guidance_scale=req.guidance_scale,
negative_prompt=req.params.negative_prompt,
height=req.params.height,
width=req.params.width,
num_frames=req.params.num_frames,
num_inference_steps=req.params.num_inference_steps,
guidance_scale=req.params.guidance_scale,
guidance_scale_2=extra.get("guidance_scale_2"),
boundary_ratio=extra.get("boundary_ratio"),
seed=req.seed,
max_sequence_length=req.max_sequence_length,
seed=req.params.seed,
max_sequence_length=req.params.max_sequence_length,
)

@nvtx_range("WanPipeline.forward")
Expand Down
22 changes: 11 additions & 11 deletions tensorrt_llm/_torch/visual_gen/models/wan/pipeline_wan_i2v.py
Original file line number Diff line number Diff line change
Expand Up @@ -393,11 +393,11 @@ def extra_param_specs(self):
def infer(self, req):
"""Run inference with request parameters."""
# Extract image from request (can be path, PIL Image, or torch.Tensor)
if req.image is None:
if req.params.image is None:
raise ValueError("I2V pipeline requires 'image' parameter")

image = req.image[0] if isinstance(req.image, list) else req.image
extra = req.extra_params or {}
image = req.params.image[0] if isinstance(req.params.image, list) else req.params.image
extra = req.params.extra_params or {}
last_image = extra.get("last_image")

if last_image is not None and isinstance(last_image, list):
Expand All @@ -406,16 +406,16 @@ def infer(self, req):
return self.forward(
image=image,
prompt=req.prompt,
negative_prompt=req.negative_prompt,
height=req.height,
width=req.width,
num_frames=req.num_frames,
num_inference_steps=req.num_inference_steps,
guidance_scale=req.guidance_scale,
negative_prompt=req.params.negative_prompt,
height=req.params.height,
width=req.params.width,
num_frames=req.params.num_frames,
num_inference_steps=req.params.num_inference_steps,
guidance_scale=req.params.guidance_scale,
guidance_scale_2=extra.get("guidance_scale_2"),
boundary_ratio=extra.get("boundary_ratio"),
seed=req.seed,
max_sequence_length=req.max_sequence_length,
seed=req.params.seed,
max_sequence_length=req.params.max_sequence_length,
last_image=last_image,
)

Expand Down
4 changes: 2 additions & 2 deletions tensorrt_llm/_torch/visual_gen/pipeline.py
Original file line number Diff line number Diff line change
Expand Up @@ -256,8 +256,8 @@ def extra_param_specs(self) -> Dict[str, ExtraParamSchema]:
def default_generation_params(self) -> dict:
"""Model-specific defaults for ``None`` fields in ``VisualGenParams``.

Keys should match ``DiffusionRequest`` field names. The executor
merges these into the request before calling ``infer()``.
Keys should match ``VisualGenParams`` field names. The executor
merges these into ``request.params`` before calling ``infer()``.
"""
return {}

Expand Down
Loading
Loading