Skip to content

Commit d2c069a

Browse files
- Full Flux2 TP1 H100 pipeline parity is exact.
- L40S TP2/TP4 has bounded BF16 TP drift. - H100 image artifacts were generated at 1024x1024 and look good. - Full model image conditioning/caption upsampling remains out of scope.
1 parent 0bd32e5 commit d2c069a

24 files changed

Lines changed: 3168 additions & 100 deletions
Lines changed: 102 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,102 @@
1+
# SPDX-License-Identifier: Apache-2.0
2+
"""Run full Flux2 text-to-image generation through FastVideo.
3+
4+
User story:
5+
"I have a local or HF Diffusers-format full Flux2 checkpoint and want a
6+
minimal text-to-image generation command that uses embedded guidance."
7+
"""
8+
import argparse
9+
import os
10+
from pathlib import Path
11+
12+
from fastvideo import VideoGenerator
13+
from fastvideo.api.sampling_param import SamplingParam
14+
15+
16+
def parse_args() -> argparse.Namespace:
17+
parser = argparse.ArgumentParser(description="Run full Flux2 text-to-image generation.")
18+
parser.add_argument(
19+
"--model-path",
20+
default="black-forest-labs/FLUX.2-dev",
21+
help="HF id or local diffusers-format full Flux2 weights directory.",
22+
)
23+
parser.add_argument(
24+
"--output",
25+
default="outputs/flux2/flux2.png",
26+
help="Output PNG path.",
27+
)
28+
parser.add_argument(
29+
"--prompt",
30+
default="a photo of a banana on a wooden table, studio lighting",
31+
help="Text prompt.",
32+
)
33+
parser.add_argument("--height", type=int, default=1024)
34+
parser.add_argument("--width", type=int, default=1024)
35+
parser.add_argument("--steps", type=int, default=50)
36+
parser.add_argument("--guidance-scale", type=float, default=4.0)
37+
parser.add_argument("--max-sequence-length", type=int, default=None)
38+
parser.add_argument("--seed", type=int, default=0)
39+
parser.add_argument("--num-gpus", type=int, default=1)
40+
parser.add_argument("--tp-size", type=int, default=None)
41+
parser.add_argument("--sp-size", type=int, default=None)
42+
parser.add_argument(
43+
"--backend",
44+
default=None,
45+
help="Set FASTVIDEO_ATTENTION_BACKEND, for example TORCH_SDPA.",
46+
)
47+
return parser.parse_args()
48+
49+
50+
def main() -> None:
51+
args = parse_args()
52+
if args.backend:
53+
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = args.backend
54+
55+
output = Path(args.output)
56+
output.parent.mkdir(parents=True, exist_ok=True)
57+
tp_size = args.tp_size if args.tp_size is not None else (
58+
args.num_gpus if args.num_gpus > 1 else 1
59+
)
60+
sp_size = args.sp_size if args.sp_size is not None else (
61+
1 if args.num_gpus > 1 else args.num_gpus
62+
)
63+
64+
generator = VideoGenerator.from_pretrained(
65+
args.model_path,
66+
num_gpus=args.num_gpus,
67+
tp_size=tp_size,
68+
sp_size=sp_size,
69+
workload_type="t2i",
70+
use_fsdp_inference=False,
71+
dit_cpu_offload=False,
72+
vae_cpu_offload=True,
73+
text_encoder_cpu_offload=True,
74+
pin_cpu_memory=False,
75+
override_pipeline_cls_name="Flux2Pipeline",
76+
)
77+
try:
78+
sampling = SamplingParam.from_pretrained(args.model_path)
79+
sampling.prompt = args.prompt
80+
sampling.height = args.height
81+
sampling.width = args.width
82+
sampling.num_frames = 1
83+
sampling.fps = 1
84+
sampling.num_inference_steps = args.steps
85+
sampling.guidance_scale = args.guidance_scale
86+
sampling.max_sequence_length = args.max_sequence_length
87+
sampling.seed = args.seed
88+
sampling.output_path = str(output)
89+
sampling.save_video = True
90+
sampling.return_frames = False
91+
92+
generator.generate_video(
93+
args.prompt,
94+
sampling_param=sampling,
95+
output_path=str(output),
96+
)
97+
finally:
98+
generator.shutdown()
99+
100+
101+
if __name__ == "__main__":
102+
main()

examples/inference/basic/basic_flux2_klein.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -52,6 +52,7 @@ def main() -> None:
5252
generator = VideoGenerator.from_pretrained(
5353
args.model_path,
5454
num_gpus=args.num_gpus,
55+
workload_type="t2i",
5556
use_fsdp_inference=False,
5657
dit_cpu_offload=False,
5758
vae_cpu_offload=True,

fastvideo/api/sampling_param.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -29,6 +29,10 @@ class SamplingParam:
2929
# Video inputs
3030
video_path: str | None = None
3131

32+
# Optional pre-generated diffusion latents. Used by parity/debug harnesses
33+
# and advanced callers that need deterministic latent reuse.
34+
latents: Any | None = None
35+
3236
# Action control inputs (Matrix-Game)
3337
mouse_cond: Any | None = None # Shape: (B, T, 2)
3438
keyboard_cond: Any | None = None # Shape: (B, T, K)
@@ -64,6 +68,7 @@ class SamplingParam:
6468
# Text inputs
6569
prompt: str | list[str] | None = None
6670
negative_prompt: str = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
71+
max_sequence_length: int | None = None
6772
prompt_path: str | None = None
6873
output_path: str = "outputs/"
6974
output_video_name: str | None = None

fastvideo/configs/models/encoders/__init__.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@
77
from fastvideo.configs.models.encoders.siglip import SiglipVisionConfig
88
from fastvideo.configs.models.encoders.reason1 import Reason1ArchConfig, Reason1Config
99
from fastvideo.configs.models.encoders.gemma import LTX2GemmaConfig
10+
from fastvideo.configs.models.encoders.mistral3 import Mistral3TextConfig
1011
from fastvideo.configs.models.encoders.qwen3 import Qwen3TextConfig
1112
from fastvideo.configs.models.encoders.stable_audio_conditioner import (StableAudioConditionerArchConfig,
1213
StableAudioConditionerConfig)
@@ -16,5 +17,5 @@
1617
"EncoderConfig", "TextEncoderConfig", "ImageEncoderConfig", "BaseEncoderOutput", "CLIPTextConfig",
1718
"CLIPVisionConfig", "WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config", "T5LargeConfig", "Qwen2_5_VLConfig",
1819
"Reason1ArchConfig", "Reason1Config", "LTX2GemmaConfig", "SiglipVisionConfig", "StableAudioConditionerArchConfig",
19-
"StableAudioConditionerConfig", "T5GemmaEncoderConfig", "Qwen3TextConfig"
20+
"StableAudioConditionerConfig", "T5GemmaEncoderConfig", "Qwen3TextConfig", "Mistral3TextConfig"
2021
]
Lines changed: 36 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,36 @@
1+
# SPDX-License-Identifier: Apache-2.0
2+
"""Mistral3 text encoder configuration for full Flux2."""
3+
from dataclasses import dataclass, field
4+
5+
from fastvideo.configs.models.encoders.base import (
6+
TextEncoderArchConfig,
7+
TextEncoderConfig,
8+
)
9+
10+
11+
@dataclass
12+
class Mistral3TextArchConfig(TextEncoderArchConfig):
13+
"""Architecture config for the Mistral3 text encoder used by full Flux2."""
14+
15+
architectures: list[str] = field(default_factory=lambda: ["Mistral3ForConditionalGeneration"])
16+
hidden_size: int = 5120
17+
num_hidden_layers: int = 40
18+
text_len: int = 512
19+
output_hidden_states: bool = True
20+
21+
def __post_init__(self) -> None:
22+
self.tokenizer_kwargs = {
23+
"padding": "max_length",
24+
"truncation": True,
25+
"max_length": self.text_len,
26+
"return_tensors": "pt",
27+
}
28+
29+
30+
@dataclass
31+
class Mistral3TextConfig(TextEncoderConfig):
32+
"""Top-level config for the Mistral3 full Flux2 text encoder."""
33+
34+
arch_config: TextEncoderArchConfig = field(default_factory=Mistral3TextArchConfig)
35+
prefix: str = "mistral3"
36+
is_chat_model: bool = True

fastvideo/configs/pipelines/flux_2.py

Lines changed: 23 additions & 32 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,7 @@
99
from fastvideo.configs.models.dits.flux_2 import Flux2Config
1010
from fastvideo.configs.models.encoders import BaseEncoderOutput
1111
from fastvideo.configs.models.encoders.base import EncoderArchConfig
12+
from fastvideo.configs.models.encoders.mistral3 import Mistral3TextConfig
1213
from fastvideo.configs.models.encoders.qwen3 import Qwen3TextConfig
1314
from fastvideo.configs.models.vaes.flux2vae import Flux2VAEConfig
1415
from fastvideo.configs.pipelines.base import PipelineConfig, preprocess_text
@@ -17,37 +18,34 @@
1718
@dataclass
1819
class Flux2PipelineConfig(PipelineConfig):
1920
"""Configuration for Flux2 image generation pipeline."""
20-
21+
2122
# Flux2-specific parameters
22-
embedded_cfg_scale: float = 4.0
23-
23+
embedded_cfg_scale: float | None = 4.0
24+
flux2_text_encoder_type: str = "mistral3"
25+
text_encoder_out_layers: tuple[int, ...] = (10, 20, 30)
26+
2427
# DiT configuration
2528
dit_config: DiTConfig = field(default_factory=Flux2Config)
2629
dit_precision: str = "bf16"
27-
30+
2831
# VAE configuration
2932
vae_config: VAEConfig = field(default_factory=Flux2VAEConfig)
3033
vae_precision: str = "fp32"
3134
vae_tiling: bool = False # Flux2 is image model, disable tiling by default
3235
vae_sp: bool = False
33-
34-
# Text encoder configuration (Flux2 uses Mistral/Qwen)
35-
text_encoder_configs: tuple[EncoderConfig, ...] = field(
36-
default_factory=lambda: (EncoderConfig(),)
37-
)
38-
text_encoder_precisions: tuple[str, ...] = field(
39-
default_factory=lambda: ("bf16",)
40-
)
41-
36+
37+
# Text encoder configuration (full Flux2 uses Mistral3)
38+
text_encoder_configs: tuple[EncoderConfig, ...] = field(default_factory=lambda: (Mistral3TextConfig(), ))
39+
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", ))
40+
4241
# Default postprocess function (can be overridden)
4342
@staticmethod
4443
def default_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
4544
"""Default text postprocessing for Flux2."""
4645
return outputs.last_hidden_state
47-
48-
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor], ...] = field(
49-
default_factory=lambda: (Flux2PipelineConfig.default_postprocess_text,)
50-
)
46+
47+
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
48+
...] = field(default_factory=lambda: (Flux2PipelineConfig.default_postprocess_text, ))
5149

5250

5351
def flux2_klein_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
@@ -57,9 +55,7 @@ def flux2_klein_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
5755
raise ValueError("Flux2 Klein requires output_hidden_states=True from text encoder")
5856
out = torch.stack([outputs.hidden_states[k] for k in hidden_states_layers], dim=1)
5957
batch_size, num_channels, seq_len, hidden_dim = out.shape
60-
prompt_embeds = out.permute(0, 2, 1, 3).reshape(
61-
batch_size, seq_len, num_channels * hidden_dim
62-
)
58+
prompt_embeds = out.permute(0, 2, 1, 3).reshape(batch_size, seq_len, num_channels * hidden_dim)
6359
return prompt_embeds
6460

6561

@@ -79,15 +75,10 @@ class Flux2KleinTextEncoderConfig(EncoderConfig):
7975
class Flux2KleinPipelineConfig(Flux2PipelineConfig):
8076
"""Configuration for Flux2 Klein (distilled, 4-step, no guidance)."""
8177
embedded_cfg_scale: float | None = None # Klein distilled: no guidance embedding (matches Diffusers)
82-
text_encoder_configs: tuple[EncoderConfig, ...] = field(
83-
default_factory=lambda: (Qwen3TextConfig(),)
84-
)
85-
text_encoder_precisions: tuple[str, ...] = field(
86-
default_factory=lambda: ("bf16",)
87-
)
88-
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
89-
default_factory=lambda: (preprocess_text,)
90-
)
91-
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor], ...] = field(
92-
default_factory=lambda: (flux2_klein_postprocess_text,)
93-
)
78+
flux2_text_encoder_type: str = "qwen3"
79+
text_encoder_out_layers: tuple[int, ...] = (9, 18, 27)
80+
text_encoder_configs: tuple[EncoderConfig, ...] = field(default_factory=lambda: (Qwen3TextConfig(), ))
81+
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", ))
82+
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(default_factory=lambda: (preprocess_text, ))
83+
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
84+
...] = field(default_factory=lambda: (flux2_klein_postprocess_text, ))

fastvideo/entrypoints/video_generator.py

Lines changed: 6 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -383,6 +383,7 @@ def _generate_video_impl(
383383
if grid_sizes is not None:
384384
kwargs['grid_sizes'] = grid_sizes
385385

386+
prompt_embeds = kwargs.pop("prompt_embeds", None)
386387
sampling_param.update(kwargs)
387388

388389
if fastvideo_args.prompt_txt is not None or sampling_param.prompt_path is not None:
@@ -432,6 +433,8 @@ def _generate_video_impl(
432433
raise ValueError("Either prompt or prompt_txt must be provided")
433434
output_path = self._prepare_output_path(sampling_param.output_path, prompt)
434435
kwargs["output_path"] = output_path
436+
if prompt_embeds is not None:
437+
kwargs["prompt_embeds"] = prompt_embeds
435438
return self._generate_single_video(
436439
prompt=prompt,
437440
sampling_param=sampling_param,
@@ -534,6 +537,7 @@ def _generate_single_video(
534537
prompt = prompt.strip()
535538
sampling_param = deepcopy(sampling_param)
536539
output_path = kwargs["output_path"]
540+
prompt_embeds = kwargs.get("prompt_embeds")
537541
sampling_param.prompt = prompt
538542
# Process negative prompt
539543
if sampling_param.negative_prompt is not None:
@@ -581,12 +585,8 @@ def _generate_single_video(
581585
VSA_sparsity=fastvideo_args.VSA_sparsity,
582586
)
583587
# Allow precomputed prompt_embeds (e.g. from diffusers) to skip text encoding
584-
if kwargs.get("prompt_embeds") is not None:
585-
batch.prompt_embeds = (
586-
list(kwargs["prompt_embeds"])
587-
if isinstance(kwargs["prompt_embeds"], (list, tuple))
588-
else [kwargs["prompt_embeds"]]
589-
)
588+
if prompt_embeds is not None:
589+
batch.prompt_embeds = (list(prompt_embeds) if isinstance(prompt_embeds, list | tuple) else [prompt_embeds])
590590

591591
# Run inference
592592
start_time = time.perf_counter()

fastvideo/models/dits/flux_2.py

Lines changed: 18 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -988,6 +988,8 @@ def forward(
988988
hidden_states: torch.Tensor,
989989
encoder_hidden_states: torch.Tensor = None,
990990
timestep: torch.LongTensor = None,
991+
img_ids: torch.Tensor = None,
992+
txt_ids: torch.Tensor = None,
991993
guidance: torch.Tensor = None,
992994
freqs_cis: torch.Tensor = None,
993995
joint_attention_kwargs: Optional[Dict[str, Any]] = None,
@@ -1047,21 +1049,22 @@ def forward(
10471049

10481050
# 3. Compute RoPE positional embeddings when not provided externally
10491051
if freqs_cis is None:
1050-
if input_was_5d:
1051-
img_h, img_w = h, w
1052-
else:
1053-
img_seq_len = hidden_states.shape[1]
1054-
img_h = img_w = int(img_seq_len ** 0.5)
1055-
1056-
txt_len = encoder_hidden_states.shape[1]
1057-
txt_ids = torch.cartesian_prod(
1058-
torch.arange(1), torch.arange(1),
1059-
torch.arange(1), torch.arange(txt_len),
1060-
).to(device=hidden_states.device)
1061-
img_ids = torch.cartesian_prod(
1062-
torch.arange(1), torch.arange(img_h),
1063-
torch.arange(img_w), torch.arange(1),
1064-
).to(device=hidden_states.device)
1052+
if txt_ids is None or img_ids is None:
1053+
if input_was_5d:
1054+
img_h, img_w = h, w
1055+
else:
1056+
img_seq_len = hidden_states.shape[1]
1057+
img_h = img_w = int(img_seq_len ** 0.5)
1058+
1059+
txt_len = encoder_hidden_states.shape[1]
1060+
txt_ids = torch.cartesian_prod(
1061+
torch.arange(1), torch.arange(1),
1062+
torch.arange(1), torch.arange(txt_len),
1063+
).to(device=hidden_states.device)
1064+
img_ids = torch.cartesian_prod(
1065+
torch.arange(1), torch.arange(img_h),
1066+
torch.arange(img_w), torch.arange(1),
1067+
).to(device=hidden_states.device)
10651068
freqs_cis = compute_flux2_freqs_cis_from_ids(
10661069
self.rotary_emb, txt_ids, img_ids, device=hidden_states.device,
10671070
)

0 commit comments

Comments
 (0)