Skip to content

[feat] Support training_cfg_rate > 0 for LTX-2 - #1752

Open
alanhuangyoo wants to merge 1 commit into
hao-ai-lab:mainfrom
alanhuangyoo:feat/ltx2-cfg-training
Open

alanhuangyoo wants to merge 1 commit into
hao-ai-lab:mainfrom
alanhuangyoo:feat/ltx2-cfg-training

Conversation

@alanhuangyoo

Copy link
Copy Markdown
Contributor

Draft — the implementation and unit tests are in place, but I have one semantics question for you before this is ready, and I have not run a convergence check yet. Details at the bottom.

Problem

LTX2Model.__init__ rejects CFG dropout outright:

raise NotImplementedError("LTX2Model only supports training_cfg_rate=0. CFG dropout "
                          "zeroes the post-connector text embeddings, which is not "
                          "what LTX-2 inference uses as the unconditional input "
                          "(an empty prompt encoded through Gemma + connector). ...")

That is accurate. The drop happens in the data layer, in get_torch_tensors_from_row_dict:

if key == 'text_embedding' and (rng.random() if rng else random.random()) < cfg_rate:
    data = np.zeros(shape, dtype=np.float32)

Zeros are not what the LTX-2 sampler feeds its unconditional branch, so training that way teaches an unconditional the sampler never uses.

Approach

Do the drop in LTX2Model.prepare_batch instead, substituting the cached unconditional embedding.

  • ensure_negative_conditioning() encodes the unconditional prompt via the shared encode_negative_prompt helper. LTX2GemmaTextEncoderModel owns Embeddings1DConnector, so what comes back is already the post-connector tensor — no separate connector pass is needed.
  • _apply_cfg_dropout() picks samples with torch.rand(..., generator=generator) and swaps in that embedding and its mask. It is a no-op when training_cfg_rate == 0, and it clones rather than mutating the batch in place.
  • set_requires_negative_conditioning(cfg_rate > 0.0) keeps the existing "don't pay 23GB of Gemma per rank for an unused embedding" property — nothing is loaded unless CFG dropout is on.
  • ltx2_training_pipeline.py passes cfg_rate=0.0 to the parquet loader so the drop happens once, in the model. This also gives the precomputed .pt path CFG dropout, which it never had (ltx2_precomputed_dataset.py has no cfg_rate).

Question before this leaves draft

What should the unconditional prompt be? I used "", following the wording in the NotImplementedError, exposed as LTX2_UNCONDITIONAL_PROMPT.

But the presets disagree with each other — ltx2_base and ltx2_distilled set negative_prompt to _LTX2_NEGATIVE_PROMPT, while ltx2_two_stage and ltx2_3_base set it to "". So "the unconditional LTX-2 samples with" is not one fixed string. Options as I see them:

  1. Always "" (what this PR does).
  2. Read SamplingParam.from_pretrained(model_path).negative_prompt, the way WanModel.ensure_negative_conditioning does — then training matches whichever preset the checkpoint ships.
  3. Make it a training config field.

I did not want to guess, since option 2 changes what gets trained depending on the checkpoint.

Verified

New unit tests in fastvideo/tests/train/models/test_ltx2_cfg_dropout.py (3 passed) cover the rate-0 no-op, that dropped samples carry the unconditional embedding and specifically not zeros, and that the caller's tensors are not mutated. They construct the model shell directly, so nothing large is loaded.

pre-commit run --files ... is clean on all three files (yapf, ruff, codespell, mypy).

pytest fastvideo/tests/train/ -k ltx2 gives the same 4 failures with and without this patch — test_ltx2_finetune_single_train_step[ltx2|ltx2_3] and test_ltx2_model_loads_and_forwards[ltx2|ltx2_3], all Could not find model at FastVideo/LTX-2.3-Distilled-Diffusers and failed to download from HF Hub. Pre-existing, environment-side.

Not verified

No convergence run. I can put a short LoRA finetune with training_cfg_rate=0.1 against cfg_rate=0 on H20 once the prompt question above is settled, so the run measures the intended behaviour.

Env: torch 2.12.0+cu130, H20, single node.

@mergify mergify Bot added type: feat New feature or capability scope: training Training pipeline, methods, configs scope: infra CI, tests, Docker, build labels Aug 24, 2026
@mergify

mergify Bot commented Aug 24, 2026 •

Copy link
Copy Markdown
Contributor

Merge Protections

🔴 1 of 1 protections blocking · waiting on 👀 reviews and 🤖 CI

Protection Waiting on
🔴 PR merge requirements 👀 reviews and 🤖 CI

🔴 PR merge requirements

Waiting for

  • #approved-reviews-by>=1
  • check-success=full-suite-passed
This rule is failing.
  • #approved-reviews-by>=1
  • check-success=full-suite-passed
  • check-success=fastcheck-passed
  • check-success~=pre-commit
  • title~=(?i)^\[(feat|feature|bugfix|fix|refactor|perf|ci|doc|docs|misc|chore|kernel|new.?model|skill|skills|infra)\]

@alanhuangyoo
alanhuangyoo marked this pull request as ready for review August 26, 2026 06:37
CFG dropout was rejected outright because the shared parquet path drops
text conditioning by zeroing the stored embedding, which is not the
unconditional input LTX-2 samples with. Do the drop inside
LTX2Model.prepare_batch so dropped samples carry the empty prompt encoded
through Gemma and the Embeddings1D connector.

LTX2GemmaTextEncoderModel owns the connector, so encode_negative_prompt
already returns the post-connector tensor. Gemma is only loaded when
training_cfg_rate > 0.
@alanhuangyoo
alanhuangyoo force-pushed the feat/ltx2-cfg-training branch from 6651cc2 to 6236e29 Compare August 26, 2026 08:12
@alanhuangyoo

Copy link
Copy Markdown
Contributor Author

@SolitaryThinker Thanks for merging #1791! Could you also take a look at this one and #1751 when you get a chance?

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

scope: infra CI, tests, Docker, build scope: training Training pipeline, methods, configs type: feat New feature or capability

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant