[diffusion] fix: improve LTX2.3 reference accuracy controls (#24022)

This commit is contained in:
Mick
2026-04-29 21:39:27 +08:00
committed by GitHub
parent 1279ae0787
commit e5c2a9b6cf
8 changed files with 98 additions and 37 deletions
@@ -64,6 +64,7 @@ class LTX23SamplingParams(LTX2SamplingParams):
audio_modality_scale: float = 3.0 audio_modality_scale: float = 3.0
audio_skip_step: int = 0 audio_skip_step: int = 0
audio_stg_blocks: list[int] = field(default_factory=lambda: [28]) audio_stg_blocks: list[int] = field(default_factory=lambda: [28])
skip_v2a_cross_attn_for_video_gt: bool = False
def build_request_extra(self) -> dict[str, Any]: def build_request_extra(self) -> dict[str, Any]:
extra = super().build_request_extra() extra = super().build_request_extra()
@@ -81,6 +82,8 @@ class LTX23SamplingParams(LTX2SamplingParams):
"audio_skip_step": self.audio_skip_step, "audio_skip_step": self.audio_skip_step,
"audio_stg_blocks": self.audio_stg_blocks, "audio_stg_blocks": self.audio_stg_blocks,
} }
if self.skip_v2a_cross_attn_for_video_gt:
extra["ltx2_skip_v2a_cross_attn_for_video_gt"] = True
return extra return extra
@@ -19,6 +19,7 @@ def update_config_from_args(
# Handle top-level attributes (no prefix) # Handle top-level attributes (no prefix)
args_not_to_remove = [ args_not_to_remove = [
"model_path", "model_path",
"disable_autocast",
] ]
args_to_remove = [] args_to_remove = []
if prefix.strip() == "": if prefix.strip() == "":
@@ -1450,6 +1450,7 @@ class LTX2VideoTransformer3DModel(CachableDiT, OffloadableDiTMixin):
base_num_frames=cross_attn_pos_embed_max_pos, base_num_frames=cross_attn_pos_embed_max_pos,
sampling_rate=16000, sampling_rate=16000,
hop_length=160, hop_length=160,
scale_factors=self.audio_scale_factors,
theta=float(arch.positional_embedding_theta), theta=float(arch.positional_embedding_theta),
causal_offset=causal_offset, causal_offset=causal_offset,
modality="audio", modality="audio",
@@ -1446,6 +1446,9 @@ class LTX2DenoisingStage(DenoisingStage):
and int(getattr(batch, "ltx2_num_image_tokens", 0)) > 0 and int(getattr(batch, "ltx2_num_image_tokens", 0)) > 0
) )
) )
skip_v2a_cross_attn_for_video_gt = bool(
batch.extra.get("ltx2_skip_v2a_cross_attn_for_video_gt", False)
)
def evaluate_stage1_guided_x0( def evaluate_stage1_guided_x0(
*, *,
@@ -1476,6 +1479,9 @@ class LTX2DenoisingStage(DenoisingStage):
encoder_hidden_states=encoder_hidden_states, encoder_hidden_states=encoder_hidden_states,
audio_encoder_hidden_states=audio_encoder_hidden_states, audio_encoder_hidden_states=audio_encoder_hidden_states,
encoder_attention_mask=encoder_attention_mask, encoder_attention_mask=encoder_attention_mask,
disable_v2a_cross_attn=(
skip_v2a_cross_attn_for_video_gt
),
) )
) )
v_neg, a_v_neg = step.current_model( v_neg, a_v_neg = step.current_model(
@@ -1485,6 +1491,9 @@ class LTX2DenoisingStage(DenoisingStage):
encoder_hidden_states=negative_encoder_hidden_states, encoder_hidden_states=negative_encoder_hidden_states,
audio_encoder_hidden_states=negative_audio_encoder_hidden_states, audio_encoder_hidden_states=negative_audio_encoder_hidden_states,
encoder_attention_mask=negative_encoder_attention_mask, encoder_attention_mask=negative_encoder_attention_mask,
disable_v2a_cross_attn=(
skip_v2a_cross_attn_for_video_gt
),
) )
) )
@@ -1510,6 +1519,9 @@ class LTX2DenoisingStage(DenoisingStage):
skip_audio_self_attn_blocks=tuple( skip_audio_self_attn_blocks=tuple(
stage1_guider_params["audio_stg_blocks"] stage1_guider_params["audio_stg_blocks"]
), ),
disable_v2a_cross_attn=(
skip_v2a_cross_attn_for_video_gt
),
) )
) )
v_ptb = v_ptb.float() v_ptb = v_ptb.float()
@@ -1539,12 +1551,14 @@ class LTX2DenoisingStage(DenoisingStage):
encoder_hidden_states=encoder_hidden_states, encoder_hidden_states=encoder_hidden_states,
audio_encoder_hidden_states=audio_encoder_hidden_states, audio_encoder_hidden_states=audio_encoder_hidden_states,
encoder_attention_mask=encoder_attention_mask, encoder_attention_mask=encoder_attention_mask,
disable_v2a_cross_attn=skip_v2a_cross_attn_for_video_gt,
), ),
LTX2GuidancePassSpec( LTX2GuidancePassSpec(
name="neg", name="neg",
encoder_hidden_states=negative_encoder_hidden_states, encoder_hidden_states=negative_encoder_hidden_states,
audio_encoder_hidden_states=negative_audio_encoder_hidden_states, audio_encoder_hidden_states=negative_audio_encoder_hidden_states,
encoder_attention_mask=negative_encoder_attention_mask, encoder_attention_mask=negative_encoder_attention_mask,
disable_v2a_cross_attn=skip_v2a_cross_attn_for_video_gt,
), ),
] ]
if need_perturbed: if need_perturbed:
@@ -1560,6 +1574,7 @@ class LTX2DenoisingStage(DenoisingStage):
skip_audio_self_attn_blocks=tuple( skip_audio_self_attn_blocks=tuple(
stage1_guider_params["audio_stg_blocks"] stage1_guider_params["audio_stg_blocks"]
), ),
disable_v2a_cross_attn=skip_v2a_cross_attn_for_video_gt,
) )
) )
if need_modality: if need_modality:
@@ -45,47 +45,59 @@ class LTX2TextConnectorStage(PipelineStage):
else None else None
) )
# Handle CFG: Concatenate negative and positive inputs if prompt_embeds is None or prompt_attention_mask is None:
if batch.do_classifier_free_guidance: raise ValueError(
# Concatenate: [Negative, Positive] "LTX2TextConnectorStage requires prompt embeddings and "
prompt_embeds = torch.cat([neg_prompt_embeds, prompt_embeds], dim=0) "attention mask."
prompt_attention_mask = torch.cat(
[neg_prompt_attention_mask, prompt_attention_mask], dim=0
) )
# Prepare additive mask for connectors (as per Diffusers implementation) if batch.do_classifier_free_guidance:
dtype = prompt_embeds.dtype if neg_prompt_embeds is None or neg_prompt_attention_mask is None:
raise ValueError(
additive_attention_mask = (prompt_attention_mask.to(torch.int64) - 1).to( "LTX2TextConnectorStage requires negative prompt embeddings "
dtype "and attention mask when classifier-free guidance is enabled."
) * torch.finfo(dtype).max
# Call connectors
# Expects: prompt_embeds, attention_mask, additive_mask=True
with set_forward_context(current_timestep=None, attn_metadata=None):
connector_prompt_embeds, connector_audio_prompt_embeds, connector_mask = (
self.connectors(
prompt_embeds, additive_attention_mask, additive_mask=True
) )
)
# Split results if CFG was enabled # Official LTX-2.3 processes positive and negative prompts through
if batch.do_classifier_free_guidance: # the connector independently; batching shifts output numerics.
neg_embeds, pos_embeds = connector_prompt_embeds.chunk(2, dim=0) dtype = prompt_embeds.dtype
neg_audio_embeds, pos_audio_embeds = connector_audio_prompt_embeds.chunk( pos_additive_mask = (prompt_attention_mask.to(torch.int64) - 1).to(
2, dim=0 dtype
) ) * torch.finfo(dtype).max
neg_mask, pos_mask = connector_mask.chunk(2, dim=0) neg_additive_mask = (neg_prompt_attention_mask.to(torch.int64) - 1).to(
dtype
) * torch.finfo(dtype).max
with set_forward_context(current_timestep=None, attn_metadata=None):
pos_embeds, pos_audio_embeds, pos_mask = self.connectors(
prompt_embeds, pos_additive_mask, additive_mask=True
)
neg_embeds, neg_audio_embeds, neg_mask = self.connectors(
neg_prompt_embeds, neg_additive_mask, additive_mask=True
)
batch.prompt_embeds = [pos_embeds] batch.prompt_embeds = [pos_embeds]
batch.audio_prompt_embeds = [pos_audio_embeds] batch.audio_prompt_embeds = [pos_audio_embeds]
batch.prompt_attention_mask = pos_mask batch.prompt_attention_mask = pos_mask
batch.negative_prompt_embeds = [neg_embeds] batch.negative_prompt_embeds = [neg_embeds]
batch.negative_audio_prompt_embeds = [neg_audio_embeds] batch.negative_audio_prompt_embeds = [neg_audio_embeds]
batch.negative_attention_mask = neg_mask batch.negative_attention_mask = neg_mask
else: else:
# Update positive fields # Prepare additive mask for connectors (as per diffusers implementation)
dtype = prompt_embeds.dtype
additive_attention_mask = (prompt_attention_mask.to(torch.int64) - 1).to(
dtype
) * torch.finfo(dtype).max
with set_forward_context(current_timestep=None, attn_metadata=None):
(
connector_prompt_embeds,
connector_audio_prompt_embeds,
connector_mask,
) = self.connectors(
prompt_embeds, additive_attention_mask, additive_mask=True
)
batch.prompt_embeds = [connector_prompt_embeds] batch.prompt_embeds = [connector_prompt_embeds]
batch.audio_prompt_embeds = [connector_audio_prompt_embeds] batch.audio_prompt_embeds = [connector_audio_prompt_embeds]
batch.prompt_attention_mask = connector_mask batch.prompt_attention_mask = connector_mask
@@ -104,16 +104,16 @@
"mean_abs_diff_threshold": 8.0 "mean_abs_diff_threshold": 8.0
}, },
"ltx_2.3_one_stage_ti2v": { "ltx_2.3_one_stage_ti2v": {
"clip_threshold": 0.57, "clip_threshold": 0.64,
"ssim_threshold": 0.42, "ssim_threshold": 0.42,
"psnr_threshold": 9.0, "psnr_threshold": 8.8,
"mean_abs_diff_threshold": 57.0 "mean_abs_diff_threshold": 59.0
}, },
"ltx_2.3_two_stage_t2v_2gpus": { "ltx_2.3_two_stage_t2v_2gpus": {
"clip_threshold": 0.78, "clip_threshold": 0.80,
"ssim_threshold": 0.20, "ssim_threshold": 0.12,
"psnr_threshold": 12.7, "psnr_threshold": 12.2,
"mean_abs_diff_threshold": 49.0 "mean_abs_diff_threshold": 50.5
}, },
"wan2_1_t2v_1.3b_teacache_enabled": { "wan2_1_t2v_1.3b_teacache_enabled": {
"clip_threshold": 0.93, "clip_threshold": 0.93,
@@ -246,6 +246,18 @@
"ssim_threshold": 0.89, "ssim_threshold": 0.89,
"psnr_threshold": 24.0, "psnr_threshold": 24.0,
"mean_abs_diff_threshold": 10.0 "mean_abs_diff_threshold": 10.0
},
"ltx_2_3_hq_pipeline": {
"clip_threshold": 0.78,
"ssim_threshold": 0.48,
"psnr_threshold": 12.0,
"mean_abs_diff_threshold": 48.0
},
"ltx_2_3_two_stage_ti2v_2gpus": {
"clip_threshold": 0.92,
"ssim_threshold": 0.58,
"psnr_threshold": 17.5,
"mean_abs_diff_threshold": 20.0
} }
}, },
"default_clip_threshold_image": 0.92, "default_clip_threshold_image": 0.92,
@@ -2576,7 +2576,7 @@
"LTX2AVLatentPreparationStage": 0.33, "LTX2AVLatentPreparationStage": 0.33,
"LTX2ImageEncodingStage": 0.04, "LTX2ImageEncodingStage": 0.04,
"LTX2AVDenoisingStage": 25873.2, "LTX2AVDenoisingStage": 25873.2,
"LTX2UpsampleStage": 216.72, "LTX2UpsampleStage": 752.07,
"LTX2RefinementStage": 2553.3, "LTX2RefinementStage": 2553.3,
"LTX2AVDecodingStage": 585.0, "LTX2AVDecodingStage": 585.0,
"per_frame_generation": null "per_frame_generation": null
@@ -210,6 +210,23 @@ class TestPipelineResolutionCliOverride(unittest.TestCase):
self.assertEqual(server_args.pipeline_config.resolution, 768) self.assertEqual(server_args.pipeline_config.resolution, 768)
def test_disable_autocast_is_preserved_after_pipeline_config_resolution(self):
parser = FlexibleArgumentParser()
ServerArgs.add_cli_args(parser)
argv = [
"--model-path",
"Qwen/Qwen-Image-Layered",
"--disable-autocast",
"true",
]
with patch.object(sys, "argv", ["sglang"] + argv):
args, unknown_args = parser.parse_known_args(argv)
server_args = ServerArgs.from_cli_args(args, unknown_args)
self.assertTrue(server_args.pipeline_config.disable_autocast)
self.assertTrue(server_args.disable_autocast)
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()