Compare commits
16
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4d8713f6e4 | ||
|
|
fda43fc610 | ||
|
|
5a8329ddf6 | ||
|
|
fcb5b465c5 | ||
|
|
4f13fb0fee | ||
|
|
d96ec99b4b | ||
|
|
2555f25cce | ||
|
|
fbe56fce8d | ||
|
|
5064bcdc47 | ||
|
|
2aa824b615 | ||
|
|
f934efc58d | ||
|
|
76b0550c15 | ||
|
|
384c1e9493 | ||
|
|
b1dbcc93f6 | ||
|
|
b93833772e | ||
|
|
30b523edd6 |
@@ -69,7 +69,7 @@ approval, then upload reviewed accepted baseline records.
|
||||
| `model_id` | Yes | Benchmark id, e.g. `wan-t2v-1.3b-2gpu`. This maps to the HF subdirectory after `sanitize(model_id)`. |
|
||||
| `gpu_type` | Yes | Exact GPU device string from the performance record, e.g. the L40S device name emitted by CI. Baselines are GPU-specific. |
|
||||
| `source_results` | Yes | One or more local paths or Buildkite artifact URLs for accepted shifted performance JSONs. Prefer normalized `normalized_perf_*.json` artifacts emitted by `compare_baseline.py`. Accept `source_result` as an alias only for a single JSON. |
|
||||
| `max_intra_batch_regression` | No | Maximum allowed regression of any source JSON against the source batch median. Default: `PERF_MAX_REGRESSION` if set, otherwise `0.05` (5%). |
|
||||
| `max_intra_batch_regression` | No | Maximum allowed regression of any source JSON against the source batch median. Default: `0.05` (5%). |
|
||||
| `intent_rationale` | Yes | One-line explanation for why the baseline shift is legitimate. This is written into provenance and should be reused in the PR. |
|
||||
|
||||
Hardcoded defaults:
|
||||
@@ -148,8 +148,7 @@ For each metric with at least two non-null source values:
|
||||
4. Stop if any source record regresses against the batch median by more than
|
||||
`max_intra_batch_regression`.
|
||||
|
||||
Default `max_intra_batch_regression` to `PERF_MAX_REGRESSION` when set,
|
||||
otherwise `0.05`. Print a table with per-source values, batch median, and
|
||||
Default `max_intra_batch_regression` to `0.05`. Print a table with per-source values, batch median, and
|
||||
worst intra-batch regression.
|
||||
|
||||
This check prevents uploading a mixed batch where one JSON is materially
|
||||
@@ -183,7 +182,7 @@ present, that run is not a valid source for baseline reseeding.
|
||||
|
||||
### 2. Sync and back up existing HF records under /tmp
|
||||
|
||||
Use `fastvideo/tests/performance/hf_store.py` helpers directly. Do **not** use
|
||||
Use `fastvideo/performance/hf_store.py` helpers directly. Do **not** use
|
||||
`compare_baseline.py` as a sync shortcut; on full main runs it can persist
|
||||
records, while this step must only fetch and back up existing history.
|
||||
|
||||
@@ -192,7 +191,7 @@ The sync command pattern is:
|
||||
```bash
|
||||
export PERFORMANCE_TRACKING_ROOT="${PERFORMANCE_TRACKING_ROOT:-/tmp/perf-tracking}"
|
||||
export HF_REPO_ID="${HF_REPO_ID:-FastVideo/performance-tracking}"
|
||||
PYTHONPATH=fastvideo/tests/performance python -c 'from hf_store import sync_from_hf; import os; sync_from_hf(os.environ["PERFORMANCE_TRACKING_ROOT"], strict=True)'
|
||||
python -c 'from fastvideo.performance.hf_store import sync_from_hf; import os; sync_from_hf(os.environ["PERFORMANCE_TRACKING_ROOT"], strict=True)'
|
||||
```
|
||||
|
||||
Then back up only the sanitized model directory under `/tmp`:
|
||||
@@ -200,8 +199,8 @@ Then back up only the sanitized model directory under `/tmp`:
|
||||
```bash
|
||||
SHORT_COMMIT=$(git rev-parse --short=12 HEAD)
|
||||
TIMESTAMP=$(date -u +%Y%m%d_%H%M%S)
|
||||
MODEL_SAFE=$(PYTHONPATH=fastvideo/tests/performance python - <<'PY'
|
||||
from hf_store import sanitize
|
||||
MODEL_SAFE=$(python - <<'PY'
|
||||
from fastvideo.performance.hf_store import sanitize
|
||||
print(sanitize("<model_id>"))
|
||||
PY
|
||||
)
|
||||
@@ -235,7 +234,7 @@ first baseline seed. Continue, but report that baseline history was empty.
|
||||
Load the last 5 successful records for the target:
|
||||
|
||||
```python
|
||||
from hf_store import load_records_for_model
|
||||
from fastvideo.performance.hf_store import load_records_for_model
|
||||
|
||||
records = load_records_for_model(
|
||||
"/tmp/perf-tracking",
|
||||
@@ -372,7 +371,7 @@ prepared records plus backup on disk.
|
||||
Use the shared storage helper so the path and repo type match CI:
|
||||
|
||||
```python
|
||||
from hf_store import upload_record
|
||||
from fastvideo.performance.hf_store import upload_record
|
||||
|
||||
upload_record("<local_record_path>", record, strict=True)
|
||||
```
|
||||
@@ -460,7 +459,7 @@ directories created for this reseed. Never remove unrelated `/tmp` contents.
|
||||
intentional baseline replacement.
|
||||
- `fastvideo/tests/performance/compare_baseline.py` — normalization, rolling
|
||||
median comparison, and persistence rules.
|
||||
- `fastvideo/tests/performance/hf_store.py` — HF sync, record loading,
|
||||
- `fastvideo/performance/hf_store.py` — HF sync, record loading,
|
||||
`sanitize()`, and `upload_record()`.
|
||||
- `fastvideo/tests/performance/test_inference_performance.py` — source result
|
||||
JSON schema.
|
||||
|
||||
@@ -1,7 +1,11 @@
|
||||
name: pre-commit
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
# pull_request_target instead of pull_request: the workflow definition and
|
||||
# the hook config are always taken from the BASE branch, so fork /
|
||||
# first-time-contributor PRs run immediately without a maintainer clicking
|
||||
# "Approve and run". The PR head is checked out as data only.
|
||||
pull_request_target:
|
||||
branches: [main]
|
||||
workflow_call:
|
||||
inputs:
|
||||
@@ -15,12 +19,25 @@ permissions:
|
||||
|
||||
jobs:
|
||||
pre-commit:
|
||||
if: github.event_name == 'workflow_call' || github.event.pull_request.draft != true
|
||||
if: github.event.pull_request.draft != true
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
with:
|
||||
ref: ${{ inputs.ref || '' }}
|
||||
# For PR events, lint the PR head — but keep the hook definitions from
|
||||
# the base branch so an untrusted PR cannot alter what gets executed.
|
||||
- name: Save trusted hook config
|
||||
if: github.event_name == 'pull_request_target'
|
||||
run: cp .pre-commit-config.yaml "$RUNNER_TEMP/trusted-pre-commit-config.yaml"
|
||||
- uses: actions/checkout@v4
|
||||
if: github.event_name == 'pull_request_target'
|
||||
with:
|
||||
ref: ${{ github.event.pull_request.head.sha }}
|
||||
persist-credentials: false
|
||||
- name: Restore trusted hook config
|
||||
if: github.event_name == 'pull_request_target'
|
||||
run: cp "$RUNNER_TEMP/trusted-pre-commit-config.yaml" .pre-commit-config.yaml
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.12"
|
||||
|
||||
@@ -112,10 +112,11 @@ per-metric policy with direction, percent threshold, absolute threshold, and a
|
||||
while it runs so pipeline stage execution times are available in
|
||||
`generate_video(...).logging_info`. Stage logs use pipeline-unique keys such as
|
||||
`prompt_encoding_stage` so duplicate stage classes do not collide. For
|
||||
`PipelineStage` entries, the extractor maps the `stage_class` field:
|
||||
`TextEncodingStage` maps to `text_encoder_time_s`, `DenoisingStage` and
|
||||
`DmdDenoisingStage` map to `dit_time_s`, and `DecodingStage` maps to
|
||||
`vae_decode_time_s`, with a fallback for older logs that used the class name as
|
||||
`PipelineStage` entries, shared component stage bases emit a stable
|
||||
`component_metric`: text encoding stages map to `text_encoder_time_s`,
|
||||
denoising stages and subclasses map to `dit_time_s`, and decoding stages map to
|
||||
`vae_decode_time_s`. The extractor falls back to known `stage_class` names for
|
||||
older logs that do not include `component_metric` or that used the class name as
|
||||
the stage key. Generator-side timings such as `PostDecodeFrameProcessStage`,
|
||||
`VideoSaveStage`, and `AudioMuxStage` are intentionally ignored. If a pipeline
|
||||
does not report one of the mapped stages, that component metric is stored as
|
||||
@@ -432,5 +433,5 @@ pipelines that did not report a mapped component stage.
|
||||
|
||||
**Component timing is `null`** — the generated result did not include a mapped
|
||||
stage in `logging_info.stages`. Check that the pipeline emits stage logging
|
||||
and that the stage name is listed in `STAGE_METRIC_MAP` in
|
||||
`test_inference_performance.py`.
|
||||
and that the stage emits `component_metric` or is covered by the legacy
|
||||
`STAGE_METRIC_MAP` fallback in `test_inference_performance.py`.
|
||||
|
||||
@@ -0,0 +1,78 @@
|
||||
# Causal Consistency Distillation: Wan 2.1 T2V 1.3B Causal
|
||||
#
|
||||
# ODE-data-free distillation. A frozen teacher takes a single CFG Euler step;
|
||||
# the student matches an EMA copy of itself at the next timestep, all under
|
||||
# clean-history teacher forcing.
|
||||
#
|
||||
# All three roles initialize from the SAME checkpoint (the teacher-forcing
|
||||
# AR-diffusion model). Point init_from at that checkpoint for a real run.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
teacher:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: false
|
||||
ema:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: false
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.consistency_model.causal_cd.CausalConsistencyDistillationMethod
|
||||
discrete_cd_N: 48
|
||||
guidance_scale: 3.0
|
||||
ema_decay: 0.99
|
||||
ema_start_step: 200
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/Wan-Syn_77x448x832_600k
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 18
|
||||
num_height: 448
|
||||
num_width: 832
|
||||
num_frames: 69
|
||||
|
||||
optimizer:
|
||||
learning_rate: 2.0e-6
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 3000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/wan2.1_causal_cd
|
||||
training_state_checkpointing_steps: 1000
|
||||
checkpoints_total_limit: 3
|
||||
|
||||
tracker:
|
||||
project_name: distillation_wan_r
|
||||
run_name: wan2.1_causal_cd_shift5
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
max_grad_norm: 1.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
@@ -0,0 +1,73 @@
|
||||
# DFSFT (Diffusion-Forcing SFT), frame-wise: Wan 2.1 T2V 1.3B Causal
|
||||
#
|
||||
# - Student: trainable causal Wan model with a block size of 1 frame
|
||||
# - Training: each frame gets its own independent noise level (frame-wise
|
||||
# diffusion forcing), versus the chunk-wise variant that shares one noise
|
||||
# level across num_frames_per_block frames.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
num_frames_per_block: 1
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.dfsft.DiffusionForcingSFTMethod
|
||||
chunk_size: 1
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/Wan-Syn_77x448x832_600k
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 18
|
||||
num_height: 448
|
||||
num_width: 832
|
||||
num_frames: 69
|
||||
|
||||
optimizer:
|
||||
learning_rate: 2.0e-6
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 4000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/wan2.1_causal_dfsft_framewise
|
||||
training_state_checkpointing_steps: 1000
|
||||
checkpoints_total_limit: 3
|
||||
|
||||
tracker:
|
||||
project_name: distillation_wan_r
|
||||
run_name: wan2.1_causal_dfsft_framewise_shift5_gauss_weight
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
pipeline_target: fastvideo.pipelines.basic.wan.wan_causal_pipeline.WanCausalPipeline
|
||||
dataset_file: examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_4.json
|
||||
every_steps: 50
|
||||
sampling_steps: [40]
|
||||
guidance_scale: 6.0
|
||||
num_frames: 69
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
@@ -0,0 +1,72 @@
|
||||
# TFSFT (Teacher-Forcing SFT): Wan 2.1 T2V 1.3B Causal
|
||||
#
|
||||
# - Student: trainable causal Wan model
|
||||
# - Training: inhomogeneous timesteps per chunk, but the causal transformer
|
||||
# denoises the current block while attending to *clean* history (clean_x),
|
||||
# not its own noisy rollout.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.tfsft.TeacherForcingSFTMethod
|
||||
chunk_size: 3
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/Wan-Syn_77x448x832_600k
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 18
|
||||
num_height: 448
|
||||
num_width: 832
|
||||
num_frames: 69
|
||||
|
||||
optimizer:
|
||||
learning_rate: 2.0e-6
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 4000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/wan2.1_causal_tfsft
|
||||
training_state_checkpointing_steps: 1000
|
||||
checkpoints_total_limit: 3
|
||||
|
||||
tracker:
|
||||
project_name: distillation_wan_r
|
||||
run_name: wan2.1_causal_tfsft_shift5_gauss_weight
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
pipeline_target: fastvideo.pipelines.basic.wan.wan_causal_pipeline.WanCausalPipeline
|
||||
dataset_file: examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_4.json
|
||||
every_steps: 50
|
||||
sampling_steps: [40]
|
||||
guidance_scale: 6.0
|
||||
num_frames: 69
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
@@ -437,6 +437,7 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
|
||||
# Causal-specific
|
||||
self.block_mask = None
|
||||
self.teacher_forcing_block_mask = None
|
||||
self.num_frame_per_block = config.arch_config.num_frames_per_block
|
||||
assert self.num_frame_per_block <= 3
|
||||
self.independent_first_frame = False
|
||||
@@ -500,6 +501,70 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
|
||||
return block_mask
|
||||
|
||||
@staticmethod
|
||||
def _prepare_teacher_forcing_mask(
|
||||
device: torch.device | str, num_frames: int = 21,
|
||||
frame_seqlen: int = 1560, num_frame_per_block=1, local_attn_size=-1
|
||||
) -> BlockMask:
|
||||
"""Attention mask for the teacher-forcing ``[clean | noisy]`` sequence.
|
||||
|
||||
A noisy token attends to its own block plus the clean context of all
|
||||
strictly previous blocks; clean tokens are block-wise causal.
|
||||
"""
|
||||
if local_attn_size != -1:
|
||||
raise NotImplementedError(
|
||||
f"Teacher forcing ignores local_attn_size={local_attn_size}: "
|
||||
"unlike the block-wise causal mask, this mask always attends "
|
||||
"to the full clean context. Windowed teacher forcing is not "
|
||||
"implemented; use local_attn_size=-1 for teacher-forcing "
|
||||
"training.")
|
||||
total_length = num_frames * frame_seqlen * 2
|
||||
padded_length = math.ceil(total_length / 128) * 128 - total_length
|
||||
|
||||
clean_ends = num_frames * frame_seqlen
|
||||
context_ends = torch.zeros(total_length + padded_length, device=device, dtype=torch.long)
|
||||
noise_context_starts = torch.zeros(total_length + padded_length, device=device, dtype=torch.long)
|
||||
noise_context_ends = torch.zeros(total_length + padded_length, device=device, dtype=torch.long)
|
||||
noise_noise_starts = torch.zeros(total_length + padded_length, device=device, dtype=torch.long)
|
||||
noise_noise_ends = torch.zeros(total_length + padded_length, device=device, dtype=torch.long)
|
||||
|
||||
attention_block_size = frame_seqlen * num_frame_per_block
|
||||
frame_indices = torch.arange(
|
||||
start=0, end=num_frames * frame_seqlen,
|
||||
step=attention_block_size, device=device, dtype=torch.long
|
||||
)
|
||||
for start in frame_indices:
|
||||
context_ends[start:start + attention_block_size] = start + attention_block_size
|
||||
|
||||
noisy_image_start_list = torch.arange(
|
||||
num_frames * frame_seqlen, total_length,
|
||||
step=attention_block_size, device=device, dtype=torch.long
|
||||
)
|
||||
noisy_image_end_list = noisy_image_start_list + attention_block_size
|
||||
for block_index, (start, end) in enumerate(zip(noisy_image_start_list, noisy_image_end_list)):
|
||||
noise_noise_starts[start:end] = start
|
||||
noise_noise_ends[start:end] = end
|
||||
noise_context_ends[start:end] = block_index * attention_block_size
|
||||
|
||||
def attention_mask(b, h, q_idx, kv_idx):
|
||||
clean_mask = (q_idx < clean_ends) & (kv_idx < context_ends[q_idx])
|
||||
c1 = (kv_idx < noise_noise_ends[q_idx]) & (kv_idx >= noise_noise_starts[q_idx])
|
||||
c2 = (kv_idx < noise_context_ends[q_idx]) & (kv_idx >= noise_context_starts[q_idx])
|
||||
noise_mask = (q_idx >= clean_ends) & (c1 | c2)
|
||||
eye_mask = q_idx == kv_idx
|
||||
return eye_mask | clean_mask | noise_mask
|
||||
|
||||
block_mask = create_block_mask(
|
||||
attention_mask, B=None, H=None,
|
||||
Q_LEN=total_length + padded_length, KV_LEN=total_length + padded_length,
|
||||
_compile=False, device=device)
|
||||
|
||||
if not dist.is_initialized() or dist.get_rank() == 0:
|
||||
print(f" cache a teacher-forcing mask with block size of {num_frame_per_block} frames")
|
||||
print(block_mask)
|
||||
|
||||
return block_mask
|
||||
|
||||
def _forward_inference(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
@@ -628,9 +693,12 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
|
||||
| None = None,
|
||||
start_frame: int = 0,
|
||||
clean_x: torch.Tensor | None = None,
|
||||
aug_t: torch.Tensor | None = None,
|
||||
**kwargs) -> torch.Tensor:
|
||||
|
||||
orig_dtype = hidden_states.dtype
|
||||
teacher_forcing = clean_x is not None
|
||||
if not isinstance(encoder_hidden_states, torch.Tensor):
|
||||
encoder_hidden_states = encoder_hidden_states[0]
|
||||
if isinstance(encoder_hidden_states_image,
|
||||
@@ -663,15 +731,26 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
freqs_cis = (freqs_cos,
|
||||
freqs_sin) if freqs_cos is not None else None
|
||||
|
||||
# Construct blockwise causal attn mask
|
||||
if self.block_mask is None:
|
||||
self.block_mask = self._prepare_blockwise_causal_attn_mask(
|
||||
device=hidden_states.device,
|
||||
num_frames=num_frames,
|
||||
frame_seqlen=post_patch_height * post_patch_width,
|
||||
num_frame_per_block=self.num_frame_per_block,
|
||||
local_attn_size=self.local_attn_size
|
||||
)
|
||||
if teacher_forcing:
|
||||
if self.teacher_forcing_block_mask is None:
|
||||
self.teacher_forcing_block_mask = self._prepare_teacher_forcing_mask(
|
||||
device=hidden_states.device,
|
||||
num_frames=num_frames,
|
||||
frame_seqlen=post_patch_height * post_patch_width,
|
||||
num_frame_per_block=self.num_frame_per_block,
|
||||
local_attn_size=self.local_attn_size,
|
||||
)
|
||||
block_mask = self.teacher_forcing_block_mask
|
||||
else:
|
||||
if self.block_mask is None:
|
||||
self.block_mask = self._prepare_blockwise_causal_attn_mask(
|
||||
device=hidden_states.device,
|
||||
num_frames=num_frames,
|
||||
frame_seqlen=post_patch_height * post_patch_width,
|
||||
num_frame_per_block=self.num_frame_per_block,
|
||||
local_attn_size=self.local_attn_size
|
||||
)
|
||||
block_mask = self.block_mask
|
||||
|
||||
hidden_states = self.patch_embedding(hidden_states)
|
||||
grid_sizes = torch.stack(
|
||||
@@ -679,6 +758,7 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
hidden_states = hidden_states.flatten(2).transpose(1, 2)
|
||||
|
||||
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
|
||||
encoder_hidden_states_text = encoder_hidden_states
|
||||
|
||||
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
|
||||
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
|
||||
@@ -694,18 +774,35 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
|
||||
assert encoder_hidden_states.dtype == orig_dtype
|
||||
|
||||
if teacher_forcing:
|
||||
# Tile RoPE/modulation so clean frame i and noisy frame i share a position.
|
||||
clean_tokens = self.patch_embedding(clean_x).flatten(2).transpose(1, 2)
|
||||
hidden_states = torch.cat([clean_tokens, hidden_states], dim=1)
|
||||
if aug_t is None:
|
||||
aug_t = torch.zeros_like(timestep)
|
||||
_, timestep_proj_clean, _, _ = self.condition_embedder(
|
||||
aug_t.flatten(), encoder_hidden_states_text, None)
|
||||
timestep_proj_clean = timestep_proj_clean.unflatten(
|
||||
1, (6, self.hidden_size)).unflatten(dim=0, sizes=timestep.shape)
|
||||
timestep_proj = torch.cat([timestep_proj_clean, timestep_proj], dim=1)
|
||||
freqs_cis = (torch.cat([freqs_cos, freqs_cos], dim=0),
|
||||
torch.cat([freqs_sin, freqs_sin], dim=0))
|
||||
|
||||
# 4. Transformer blocks
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
for block in self.blocks:
|
||||
hidden_states = self._gradient_checkpointing_func(
|
||||
block, hidden_states, encoder_hidden_states,
|
||||
timestep_proj, freqs_cis,
|
||||
block_mask=self.block_mask)
|
||||
block_mask=block_mask)
|
||||
else:
|
||||
for block in self.blocks:
|
||||
hidden_states = block(hidden_states, encoder_hidden_states,
|
||||
timestep_proj, freqs_cis,
|
||||
block_mask=self.block_mask)
|
||||
block_mask=block_mask)
|
||||
|
||||
if teacher_forcing:
|
||||
hidden_states = hidden_states[:, hidden_states.shape[1] // 2:]
|
||||
|
||||
# 5. Output norm, projection & unpatchify
|
||||
temb = temb.unflatten(dim=0, sizes=timestep.shape).unsqueeze(2)
|
||||
|
||||
@@ -310,8 +310,12 @@ def load_records_for_model(
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_NUMERIC_COLS = (
|
||||
"latency", "throughput", "memory",
|
||||
"text_encoder_time_s", "dit_time_s", "vae_decode_time_s",
|
||||
"latency",
|
||||
"throughput",
|
||||
"memory",
|
||||
"text_encoder_time_s",
|
||||
"dit_time_s",
|
||||
"vae_decode_time_s",
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -36,6 +36,7 @@ DEFAULT_METRIC_POLICIES: tuple[MetricPolicy, ...] = (
|
||||
MetricPolicy("vae_decode_time_s", "VAE Decode", 3, True, 0.05, 0.25),
|
||||
)
|
||||
|
||||
|
||||
def _optional_float(value: Any) -> float | None:
|
||||
if value is None or isinstance(value, bool):
|
||||
return None
|
||||
@@ -57,9 +58,7 @@ def _optional_bool(value: Any) -> bool | None:
|
||||
return None
|
||||
|
||||
|
||||
def resolve_metric_policies(
|
||||
threshold_overrides: Mapping[str, Any] | None,
|
||||
) -> tuple[MetricPolicy, ...]:
|
||||
def resolve_metric_policies(threshold_overrides: Mapping[str, Any] | None, ) -> tuple[MetricPolicy, ...]:
|
||||
"""Return default metric policies with optional per-metric overrides."""
|
||||
|
||||
if not isinstance(threshold_overrides, Mapping):
|
||||
@@ -80,25 +79,15 @@ def resolve_metric_policies(
|
||||
label=base_policy.label,
|
||||
precision=base_policy.precision,
|
||||
lower_is_better=base_policy.lower_is_better,
|
||||
threshold_percent=(
|
||||
base_policy.threshold_percent
|
||||
if threshold_percent is None
|
||||
else threshold_percent
|
||||
),
|
||||
threshold_absolute=(
|
||||
base_policy.threshold_absolute
|
||||
if threshold_absolute is None
|
||||
else threshold_absolute
|
||||
),
|
||||
threshold_percent=(base_policy.threshold_percent if threshold_percent is None else threshold_percent),
|
||||
threshold_absolute=(base_policy.threshold_absolute
|
||||
if threshold_absolute is None else threshold_absolute),
|
||||
gated=base_policy.gated if gated is None else gated,
|
||||
)
|
||||
)
|
||||
))
|
||||
return tuple(policies)
|
||||
|
||||
|
||||
def serialize_metric_thresholds(
|
||||
policies: tuple[MetricPolicy, ...],
|
||||
) -> dict[str, dict[str, float | bool]]:
|
||||
def serialize_metric_thresholds(policies: tuple[MetricPolicy, ...], ) -> dict[str, dict[str, float | bool]]:
|
||||
return {
|
||||
policy.key: {
|
||||
"threshold_percent": policy.threshold_percent,
|
||||
@@ -118,10 +107,7 @@ def regression_delta(
|
||||
return None
|
||||
absolute_delta = current - baseline if policy.lower_is_better else baseline - current
|
||||
percent_delta = absolute_delta / baseline
|
||||
threshold_exceeded = (
|
||||
percent_delta > policy.threshold_percent
|
||||
and absolute_delta > policy.threshold_absolute
|
||||
)
|
||||
threshold_exceeded = (percent_delta > policy.threshold_percent and absolute_delta > policy.threshold_absolute)
|
||||
return MetricDelta(
|
||||
absolute=absolute_delta,
|
||||
percent=percent_delta,
|
||||
|
||||
@@ -34,6 +34,7 @@ class PipelineStage(ABC):
|
||||
composed with other stages to create a complete pipeline. Each stage is responsible
|
||||
for a specific part of the process, such as prompt encoding, latent preparation, etc.
|
||||
"""
|
||||
performance_component_metric: str | None = None
|
||||
|
||||
def verify_input(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""
|
||||
@@ -155,6 +156,9 @@ class PipelineStage(ABC):
|
||||
logger.info("[%s] Execution completed in %s ms", stage_name, execution_time * 1000)
|
||||
batch.logging_info.add_stage_execution_time(stage_key, execution_time)
|
||||
batch.logging_info.add_stage_metric(stage_key, "stage_class", stage_class_name)
|
||||
component_metric = self.performance_component_metric
|
||||
if component_metric is not None:
|
||||
batch.logging_info.add_stage_metric(stage_key, "component_metric", component_metric)
|
||||
except Exception as e:
|
||||
torch.cuda.synchronize()
|
||||
execution_time = time.perf_counter() - start_time
|
||||
|
||||
@@ -28,6 +28,7 @@ class DecodingStage(PipelineStage):
|
||||
This stage handles the decoding of latent representations into the final
|
||||
output format (e.g., pixel values).
|
||||
"""
|
||||
performance_component_metric = "vae_decode_time_s"
|
||||
|
||||
def __init__(self, vae, pipeline=None) -> None:
|
||||
self.vae: ParallelTiledVAE = vae
|
||||
|
||||
@@ -51,6 +51,7 @@ class DenoisingStage(PipelineStage):
|
||||
This stage handles the iterative denoising process that transforms
|
||||
the initial noise into the final output.
|
||||
"""
|
||||
performance_component_metric = "dit_time_s"
|
||||
|
||||
def __init__(self, transformer, scheduler, pipeline=None, transformer_2=None, vae=None) -> None:
|
||||
super().__init__()
|
||||
@@ -1187,6 +1188,7 @@ class Cosmos25V2WDenoisingStage(Cosmos25DenoisingStage):
|
||||
|
||||
class Cosmos25AutoDenoisingStage(PipelineStage):
|
||||
"""Route Cosmos 2.5 denoising to T2W vs V2W/I2W."""
|
||||
performance_component_metric = "dit_time_s"
|
||||
|
||||
def __init__(self, transformer, scheduler) -> None:
|
||||
super().__init__()
|
||||
|
||||
@@ -24,6 +24,7 @@ class TextEncodingStage(PipelineStage):
|
||||
This stage handles the encoding of text prompts into the embedding space
|
||||
expected by the diffusion model.
|
||||
"""
|
||||
performance_component_metric = "text_encoder_time_s"
|
||||
|
||||
def __init__(self, text_encoders, tokenizers) -> None:
|
||||
"""
|
||||
@@ -350,6 +351,7 @@ class Cosmos25TextEncodingStage(PipelineStage):
|
||||
Cosmos 2.5 uses Reason1 (Qwen2.5-VL) and relies on the encoder's
|
||||
`compute_text_embeddings_online()`.
|
||||
"""
|
||||
performance_component_metric = "text_encoder_time_s"
|
||||
|
||||
def __init__(self, text_encoder) -> None:
|
||||
super().__init__()
|
||||
|
||||
@@ -0,0 +1,85 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Guard: every test directory must be collected by some CI lane or be on
|
||||
the explicit allowlist below.
|
||||
|
||||
Three separate incidents on 2026-07-05 found test files that no CI lane
|
||||
ever collects (fastvideo/tests/stages/, tests/local_tests/ additions in
|
||||
PR #1509, and this sweep found seven dark directories in total): the tests
|
||||
pass review, merge, and then silently never run. This test makes going
|
||||
dark an explicit, reviewed decision instead of an accident: adding a new
|
||||
test directory fails CI until it is either wired into a lane or
|
||||
allowlisted here with a reason.
|
||||
|
||||
Pure text analysis — no fastvideo imports, no GPU, no torch.
|
||||
"""
|
||||
from pathlib import Path
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
TESTS_ROOT = REPO_ROOT / "fastvideo" / "tests"
|
||||
|
||||
# Files whose text constitutes "a CI lane references this directory".
|
||||
CI_SOURCES = [
|
||||
TESTS_ROOT / "modal" / "pr_test.py",
|
||||
TESTS_ROOT / "modal" / "ssim_test.py",
|
||||
*sorted((REPO_ROOT / ".buildkite").rglob("*.yml")),
|
||||
*sorted((REPO_ROOT / ".buildkite").rglob("*.sh")),
|
||||
]
|
||||
|
||||
# Directories that intentionally have no CI lane today. Every entry needs a
|
||||
# reason; remove the entry when the directory gets wired into a lane.
|
||||
# State as found on 2026-07-05 — these SHOULD shrink over time, not grow.
|
||||
ALLOWLIST = {
|
||||
"attention": "no lane yet — GPU attention-backend tests, run manually",
|
||||
"audio": "no lane yet — audio encoder tests, run manually",
|
||||
"distributed": "no lane yet — multi-GPU torchrun tests, run manually",
|
||||
"hooks": "no lane yet — run manually",
|
||||
"layers": "no lane yet — torchrun FSDP dispatch tests, run manually",
|
||||
"nightly": "by design: nightly cadence, not per-PR",
|
||||
"ops": "no lane yet — GPU op tests, run manually",
|
||||
"modal": "CI infrastructure itself, not a test suite",
|
||||
}
|
||||
|
||||
|
||||
def _dirs_with_tests() -> list[str]:
|
||||
dirs = []
|
||||
for child in sorted(TESTS_ROOT.iterdir()):
|
||||
if child.is_dir() and any(child.rglob("test_*.py")):
|
||||
dirs.append(child.name)
|
||||
return dirs
|
||||
|
||||
|
||||
def _ci_text() -> str:
|
||||
return "\n".join(
|
||||
src.read_text(errors="replace") for src in CI_SOURCES if src.exists())
|
||||
|
||||
|
||||
def test_every_test_directory_is_collected_or_allowlisted():
|
||||
ci_text = _ci_text()
|
||||
dark = [
|
||||
name for name in _dirs_with_tests()
|
||||
if f"tests/{name}" not in ci_text and name not in ALLOWLIST
|
||||
]
|
||||
assert not dark, (
|
||||
f"Test directories not referenced by any CI lane and not "
|
||||
f"allowlisted: {dark}. Wire them into a lane in "
|
||||
f"fastvideo/tests/modal/pr_test.py (or a Buildkite step), or add an "
|
||||
f"allowlist entry with a reason in {__file__}.")
|
||||
|
||||
|
||||
def test_local_tests_stays_out_of_ci():
|
||||
# tests/local_tests/ (repo root) is developer-local by design (author
|
||||
# decision, 2026-07-05): parity scaffolds and machine-specific checks
|
||||
# that must never gate CI. Fail if any CI source starts collecting it.
|
||||
assert "tests/local_tests" not in _ci_text(), (
|
||||
"tests/local_tests/ is local-only by design; remove the CI "
|
||||
"reference or move the tests into a fastvideo/tests/ lane.")
|
||||
|
||||
|
||||
def test_allowlist_entries_are_still_real_directories():
|
||||
# A stale allowlist hides regressions; entries must track reality.
|
||||
missing = [
|
||||
name for name in ALLOWLIST
|
||||
if name != "modal" and not (TESTS_ROOT / name).is_dir()
|
||||
]
|
||||
assert not missing, (
|
||||
f"Allowlisted directories no longer exist — remove them: {missing}")
|
||||
@@ -179,7 +179,11 @@ def _extract_component_times(result: dict) -> dict[str, float | None]:
|
||||
logger.debug("Skipping malformed stage '%s' data: %r", stage_name, stage_data)
|
||||
continue
|
||||
stage_class = stage_data.get("stage_class", stage_name)
|
||||
metric_key = STAGE_METRIC_MAP.get(stage_class)
|
||||
component_metric = stage_data.get("component_metric")
|
||||
if isinstance(component_metric, str) and component_metric in component_times:
|
||||
metric_key = component_metric
|
||||
else:
|
||||
metric_key = STAGE_METRIC_MAP.get(stage_class)
|
||||
if metric_key is None:
|
||||
logger.debug("Unmapped stage '%s' class '%s' (%.3fs)",
|
||||
stage_name,
|
||||
|
||||
@@ -1,13 +1,20 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from fastvideo.pipelines.pipeline_batch_info import PipelineLoggingInfo
|
||||
from fastvideo.pipelines.stages.denoising import Cosmos25AutoDenoisingStage, DenoisingStage
|
||||
from fastvideo.pipelines.stages.text_encoding import Cosmos25TextEncodingStage
|
||||
from fastvideo.tests.performance.test_inference_performance import _extract_component_times
|
||||
|
||||
|
||||
class SubclassStyleDenoisingStage(DenoisingStage):
|
||||
pass
|
||||
|
||||
|
||||
def test_extract_component_times_handles_pipeline_logging_info_object():
|
||||
logging_info = PipelineLoggingInfo()
|
||||
logging_info.add_stage_execution_time("prompt_encoding_stage", 1.25)
|
||||
logging_info.add_stage_metric("prompt_encoding_stage", "stage_class", "TextEncodingStage")
|
||||
logging_info.add_stage_metric("prompt_encoding_stage", "component_metric", "text_encoder_time_s")
|
||||
|
||||
assert _extract_component_times({"logging_info": logging_info}) == {
|
||||
"text_encoder_time_s": 1.25,
|
||||
@@ -45,6 +52,35 @@ def test_extract_component_times_uses_stage_class_for_pipeline_stage_keys():
|
||||
}
|
||||
|
||||
|
||||
def test_denoising_stage_subclasses_inherit_component_metric():
|
||||
assert SubclassStyleDenoisingStage.performance_component_metric == "dit_time_s"
|
||||
|
||||
|
||||
def test_cosmos25_direct_pipeline_stages_define_component_metrics():
|
||||
assert Cosmos25TextEncodingStage.performance_component_metric == "text_encoder_time_s"
|
||||
assert Cosmos25AutoDenoisingStage.performance_component_metric == "dit_time_s"
|
||||
|
||||
|
||||
def test_extract_component_times_uses_component_metric_for_stage_subclasses():
|
||||
result = {
|
||||
"logging_info": {
|
||||
"stages": {
|
||||
"denoising_stage": {
|
||||
"execution_time": 4.2,
|
||||
"stage_class": "CosmosDenoisingStage",
|
||||
"component_metric": "dit_time_s",
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
assert _extract_component_times(result) == {
|
||||
"text_encoder_time_s": None,
|
||||
"dit_time_s": 4.2,
|
||||
"vae_decode_time_s": None,
|
||||
}
|
||||
|
||||
|
||||
def test_extract_component_times_keeps_legacy_class_name_keys():
|
||||
# Backward-compatibility check for logs produced before pipeline-unique
|
||||
# stage keys carried a separate stage_class field.
|
||||
|
||||
@@ -0,0 +1,56 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Minimum config to run a single training step of
|
||||
# CausalConsistencyDistillationMethod on WanCausalModel for the
|
||||
# per-method smoke test. Uses the real Wan 2.1 1.3B checkpoint for both
|
||||
# the trainable student and the frozen teacher (AR Euler-step target),
|
||||
# with tiny synthetic latents.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
teacher:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: false
|
||||
ema:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: false
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.consistency_model.causal_cd.CausalConsistencyDistillationMethod
|
||||
discrete_cd_N: 12
|
||||
guidance_scale: 3.0
|
||||
ema_decay: 0.95
|
||||
|
||||
training:
|
||||
dit_precision: bf16
|
||||
|
||||
distributed:
|
||||
num_gpus: 1
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
|
||||
data:
|
||||
seed: 42
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
num_latent_t: 6
|
||||
num_height: 64
|
||||
num_width: 64
|
||||
num_frames: 21
|
||||
|
||||
optimizer:
|
||||
learning_rate: 2.0e-6
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 1
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
pipeline: {}
|
||||
@@ -0,0 +1,46 @@
|
||||
# Minimum config to run a single training step of
|
||||
# DiffusionForcingSFTMethod on a frame-wise WanCausalModel
|
||||
# (num_frames_per_block=1, chunk_size=1) for the per-method smoke
|
||||
# test. Uses the real Wan 2.1 1.3B checkpoint with tiny synthetic
|
||||
# latents.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
num_frames_per_block: 1
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.dfsft.DiffusionForcingSFTMethod
|
||||
chunk_size: 1
|
||||
|
||||
training:
|
||||
dit_precision: bf16
|
||||
|
||||
distributed:
|
||||
num_gpus: 1
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
|
||||
data:
|
||||
seed: 42
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
num_latent_t: 6
|
||||
num_height: 64
|
||||
num_width: 64
|
||||
num_frames: 21
|
||||
|
||||
optimizer:
|
||||
learning_rate: 2.0e-6
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 1
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
pipeline: {}
|
||||
@@ -0,0 +1,45 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Minimum config to run a single training step of
|
||||
# TeacherForcingSFTMethod on WanCausalModel for the per-method smoke
|
||||
# test. Identical to wan_causal_t2v_dfsft_min.yaml except the method,
|
||||
# which feeds clean history to the causal transformer (teacher forcing).
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.tfsft.TeacherForcingSFTMethod
|
||||
chunk_size: 3
|
||||
|
||||
training:
|
||||
dit_precision: bf16
|
||||
|
||||
distributed:
|
||||
num_gpus: 1
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
|
||||
data:
|
||||
seed: 42
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
num_latent_t: 6
|
||||
num_height: 64
|
||||
num_width: 64
|
||||
num_frames: 21
|
||||
|
||||
optimizer:
|
||||
learning_rate: 2.0e-6
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 1
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
pipeline: {}
|
||||
@@ -0,0 +1,147 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Per-method GPU smoke test: ``WanCausalModel`` + ``CausalConsistencyDistillationMethod``.
|
||||
|
||||
Mirrors ``test_wan_causal_dfsft.py``. Causal consistency distillation
|
||||
bootstraps a consistency MSE between the student's ``x0`` at ``t`` and an EMA
|
||||
copy of the student at ``t_next``, where ``t_next`` is produced online by a
|
||||
single CFG Euler step of a frozen teacher (all under clean-history teacher
|
||||
forcing). This test exercises the full step: finite loss, nonzero student
|
||||
gradients, frozen teacher, and a post-step EMA update.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
|
||||
os.environ.setdefault("MASTER_ADDR", "localhost")
|
||||
os.environ.setdefault("MASTER_PORT", "29519")
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo.train.methods.consistency_model.causal_cd import (
|
||||
CausalConsistencyDistillationMethod, )
|
||||
from fastvideo.train.models.wan import WanCausalModel
|
||||
from fastvideo.train.utils.config import load_run_config
|
||||
|
||||
|
||||
_FIXTURE = str(
|
||||
Path(__file__).resolve().parent.parent / "fixtures"
|
||||
/ "wan_causal_t2v_causal_cd_min.yaml")
|
||||
|
||||
|
||||
def _build_synthetic_batch(
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
) -> dict[str, torch.Tensor]:
|
||||
batch_size = 1
|
||||
return {
|
||||
"text_embedding":
|
||||
torch.randn(batch_size, 16, 4096, device=device, dtype=dtype),
|
||||
"text_attention_mask":
|
||||
torch.ones(batch_size, 16, device=device),
|
||||
"vae_latent":
|
||||
torch.randn(batch_size, 16, 6, 8, 8, device=device, dtype=dtype),
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("distributed_setup")
|
||||
def test_wan_causal_cd_single_train_step(
|
||||
monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("requires CUDA")
|
||||
|
||||
cfg = load_run_config(_FIXTURE)
|
||||
|
||||
device = torch.device("cuda:0")
|
||||
dtype = torch.bfloat16
|
||||
|
||||
monkeypatch.setattr(
|
||||
"fastvideo.train.utils.dataloader."
|
||||
"build_parquet_t2v_train_dataloader",
|
||||
lambda *args, **kwargs: None,
|
||||
)
|
||||
|
||||
student = WanCausalModel(
|
||||
init_from=cfg.models["student"]["init_from"],
|
||||
training_config=cfg.training,
|
||||
trainable=True,
|
||||
)
|
||||
student.transformer = student.transformer.to(device=device, dtype=dtype)
|
||||
|
||||
teacher = WanCausalModel(
|
||||
init_from=cfg.models["teacher"]["init_from"],
|
||||
training_config=cfg.training,
|
||||
trainable=False,
|
||||
)
|
||||
teacher.transformer = teacher.transformer.to(device=device, dtype=dtype)
|
||||
|
||||
ema = WanCausalModel(
|
||||
init_from=cfg.models["ema"]["init_from"],
|
||||
training_config=cfg.training,
|
||||
trainable=False,
|
||||
)
|
||||
ema.transformer = ema.transformer.to(device=device, dtype=dtype)
|
||||
|
||||
method = CausalConsistencyDistillationMethod(
|
||||
cfg=cfg,
|
||||
role_models={"student": student, "teacher": teacher, "ema": ema},
|
||||
)
|
||||
method.on_train_start()
|
||||
|
||||
batch = _build_synthetic_batch(device, dtype)
|
||||
loss_map, outputs, _metrics = method.single_train_step(batch, iteration=0)
|
||||
|
||||
loss = loss_map["total_loss"]
|
||||
assert torch.is_tensor(loss), "total_loss must be a torch.Tensor"
|
||||
assert torch.isfinite(loss).item(), (
|
||||
f"total_loss is not finite: {loss.item()}")
|
||||
|
||||
method.backward(loss_map, outputs, grad_accum_rounds=1)
|
||||
|
||||
blocks = student.transformer.blocks
|
||||
assert blocks is not None and len(blocks) > 0
|
||||
layer0 = blocks[0]
|
||||
trainable = [p for p in layer0.parameters() if p.requires_grad]
|
||||
assert len(trainable) > 0, "student layer 0 has no trainable parameters"
|
||||
for i, p in enumerate(trainable):
|
||||
assert p.grad is not None, f"student layer 0 param[{i}] has None grad"
|
||||
assert torch.isfinite(p.grad).all().item(), (
|
||||
f"student layer 0 param[{i}] grad contains NaN/Inf")
|
||||
assert any(p.grad.detach().float().norm().item() > 0.0 for p in trainable), (
|
||||
"all student layer-0 grads are exactly zero; consistency loss "
|
||||
"did not reach the first transformer block")
|
||||
|
||||
# Teacher must stay frozen.
|
||||
assert all(not p.requires_grad for p in teacher.transformer.parameters()), (
|
||||
"teacher must be frozen for Causal-CD")
|
||||
|
||||
# The EMA model and student start from the same checkpoint, so the first
|
||||
# parameter must match before any update. FSDP fully_shard params are
|
||||
# DTensors; compare the local shards (torch.equal is unsupported on
|
||||
# DTensor).
|
||||
def _local(p: torch.Tensor) -> torch.Tensor:
|
||||
return p.to_local() if hasattr(p, "to_local") else p
|
||||
|
||||
ema_param = next(ema.transformer.parameters())
|
||||
student_param = next(student.transformer.parameters())
|
||||
assert torch.equal(_local(ema_param), _local(student_param)), (
|
||||
"EMA model should start identical to the student (same checkpoint)")
|
||||
|
||||
# The EMA update must move EMA toward the student. Apply a visibly large
|
||||
# perturbation so the bf16 lerp is well above rounding noise (the real
|
||||
# optimizer step at lr=2e-6 would be sub-ULP in bf16).
|
||||
with torch.no_grad():
|
||||
student_param.add_(1.0)
|
||||
before = _local(ema_param).detach().float().clone()
|
||||
method._update_ema()
|
||||
after = _local(ema_param).detach().float()
|
||||
assert not torch.equal(before, after), (
|
||||
"EMA weights did not move after _update_ema")
|
||||
# EMA = decay*ema + (1-decay)*student moves ~ (1-decay) of the gap.
|
||||
expected = before + (1.0 - method._ema_decay) * (
|
||||
_local(student_param).detach().float() - before)
|
||||
assert torch.allclose(after, expected, atol=1e-2), (
|
||||
"EMA update did not follow the expected lerp")
|
||||
@@ -0,0 +1,109 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Per-method GPU smoke test: frame-wise ``WanCausalModel`` + ``DiffusionForcingSFTMethod``.
|
||||
|
||||
Mirrors ``test_wan_causal_dfsft.py`` but with a block size of 1 frame
|
||||
(``num_frames_per_block=1`` on the model, ``chunk_size=1`` on the method),
|
||||
so each frame gets its own independent noise level. The test asserts the
|
||||
override took effect and runs one train step: forward, finite loss, and
|
||||
nonzero gradients reaching the first transformer block.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
|
||||
os.environ.setdefault("MASTER_ADDR", "localhost")
|
||||
os.environ.setdefault("MASTER_PORT", "29520")
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo.train.methods.fine_tuning.dfsft import (
|
||||
DiffusionForcingSFTMethod, )
|
||||
from fastvideo.train.models.wan import WanCausalModel
|
||||
from fastvideo.train.utils.config import load_run_config
|
||||
|
||||
|
||||
_FIXTURE = str(
|
||||
Path(__file__).resolve().parent.parent / "fixtures"
|
||||
/ "wan_causal_t2v_dfsft_framewise_min.yaml")
|
||||
|
||||
|
||||
def _build_synthetic_batch(
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
) -> dict[str, torch.Tensor]:
|
||||
batch_size = 1
|
||||
return {
|
||||
"text_embedding":
|
||||
torch.randn(batch_size, 16, 4096, device=device, dtype=dtype),
|
||||
"text_attention_mask":
|
||||
torch.ones(batch_size, 16, device=device),
|
||||
"vae_latent":
|
||||
torch.randn(batch_size, 16, 6, 8, 8, device=device, dtype=dtype),
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("distributed_setup")
|
||||
def test_wan_causal_dfsft_framewise_single_train_step(
|
||||
monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("requires CUDA")
|
||||
|
||||
cfg = load_run_config(_FIXTURE)
|
||||
|
||||
device = torch.device("cuda:0")
|
||||
dtype = torch.bfloat16
|
||||
|
||||
monkeypatch.setattr(
|
||||
"fastvideo.train.utils.dataloader."
|
||||
"build_parquet_t2v_train_dataloader",
|
||||
lambda *args, **kwargs: None,
|
||||
)
|
||||
|
||||
student_cfg = cfg.models["student"]
|
||||
model = WanCausalModel(
|
||||
init_from=student_cfg["init_from"],
|
||||
training_config=cfg.training,
|
||||
trainable=True,
|
||||
num_frames_per_block=student_cfg.get("num_frames_per_block"),
|
||||
)
|
||||
assert model.transformer.num_frame_per_block == 1, (
|
||||
"frame-wise override did not reach the transformer")
|
||||
model.transformer = model.transformer.to(device=device, dtype=dtype)
|
||||
|
||||
method = DiffusionForcingSFTMethod(
|
||||
cfg=cfg,
|
||||
role_models={"student": model},
|
||||
)
|
||||
method.on_train_start()
|
||||
|
||||
batch = _build_synthetic_batch(device, dtype)
|
||||
loss_map, outputs, _metrics = method.single_train_step(batch, iteration=0)
|
||||
|
||||
loss = loss_map["total_loss"]
|
||||
assert torch.is_tensor(loss), "total_loss must be a torch.Tensor"
|
||||
assert torch.isfinite(loss).item(), (
|
||||
f"total_loss is not finite: {loss.item()}")
|
||||
|
||||
method.backward(loss_map, outputs, grad_accum_rounds=1)
|
||||
|
||||
blocks = getattr(model.transformer, "blocks", None)
|
||||
assert blocks is not None and len(blocks) > 0
|
||||
layer0 = blocks[0]
|
||||
|
||||
trainable = [p for p in layer0.parameters() if p.requires_grad]
|
||||
assert len(trainable) > 0, "layer 0 has no trainable parameters"
|
||||
|
||||
for i, p in enumerate(trainable):
|
||||
assert p.grad is not None, f"layer 0 param[{i}] has None grad"
|
||||
assert torch.isfinite(p.grad).all().item(), (
|
||||
f"layer 0 param[{i}] grad contains NaN/Inf")
|
||||
|
||||
any_nonzero = any(
|
||||
p.grad.detach().float().norm().item() > 0.0 for p in trainable)
|
||||
assert any_nonzero, (
|
||||
"all layer-0 grads are exactly zero; backward did not "
|
||||
"reach the first transformer block")
|
||||
@@ -0,0 +1,113 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Per-method GPU smoke test: ``WanCausalModel`` + ``TeacherForcingSFTMethod``.
|
||||
|
||||
Mirrors ``test_wan_causal_dfsft.py``. Teacher forcing concatenates a clean
|
||||
context copy of every frame inside the causal transformer (``clean_x``) and
|
||||
denoises the current block while attending to *clean* history. This test
|
||||
exercises the ``clean_x`` path end-to-end: forward, finite loss, and nonzero
|
||||
gradients reaching the first transformer block.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
|
||||
os.environ.setdefault("MASTER_ADDR", "localhost")
|
||||
os.environ.setdefault("MASTER_PORT", "29518")
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo.train.methods.fine_tuning.tfsft import (
|
||||
TeacherForcingSFTMethod, )
|
||||
from fastvideo.train.models.wan import WanCausalModel
|
||||
from fastvideo.train.utils.config import load_run_config
|
||||
|
||||
|
||||
_FIXTURE = str(
|
||||
Path(__file__).resolve().parent.parent / "fixtures"
|
||||
/ "wan_causal_t2v_tfsft_min.yaml")
|
||||
|
||||
|
||||
def _build_synthetic_batch(
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
) -> dict[str, torch.Tensor]:
|
||||
batch_size = 1
|
||||
return {
|
||||
"text_embedding":
|
||||
torch.randn(batch_size, 16, 4096, device=device, dtype=dtype),
|
||||
"text_attention_mask":
|
||||
torch.ones(batch_size, 16, device=device),
|
||||
"vae_latent":
|
||||
torch.randn(batch_size, 16, 6, 8, 8, device=device, dtype=dtype),
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("distributed_setup")
|
||||
def test_wan_causal_tfsft_single_train_step(
|
||||
monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("requires CUDA")
|
||||
|
||||
cfg = load_run_config(_FIXTURE)
|
||||
|
||||
device = torch.device("cuda:0")
|
||||
dtype = torch.bfloat16
|
||||
|
||||
monkeypatch.setattr(
|
||||
"fastvideo.train.utils.dataloader."
|
||||
"build_parquet_t2v_train_dataloader",
|
||||
lambda *args, **kwargs: None,
|
||||
)
|
||||
|
||||
model = WanCausalModel(
|
||||
init_from=cfg.models["student"]["init_from"],
|
||||
training_config=cfg.training,
|
||||
trainable=True,
|
||||
)
|
||||
model.transformer = model.transformer.to(device=device, dtype=dtype)
|
||||
|
||||
method = TeacherForcingSFTMethod(
|
||||
cfg=cfg,
|
||||
role_models={"student": model},
|
||||
)
|
||||
method.on_train_start()
|
||||
|
||||
batch = _build_synthetic_batch(device, dtype)
|
||||
loss_map, outputs, _metrics = method.single_train_step(batch, iteration=0)
|
||||
|
||||
loss = loss_map["total_loss"]
|
||||
assert torch.is_tensor(loss), "total_loss must be a torch.Tensor"
|
||||
assert torch.isfinite(loss).item(), (
|
||||
f"total_loss is not finite: {loss.item()}")
|
||||
|
||||
method.backward(loss_map, outputs, grad_accum_rounds=1)
|
||||
|
||||
blocks = getattr(model.transformer, "blocks", None)
|
||||
assert blocks is not None and len(blocks) > 0, (
|
||||
"CausalWanTransformer is expected to expose ``.blocks``")
|
||||
layer0 = blocks[0]
|
||||
|
||||
trainable = [p for p in layer0.parameters() if p.requires_grad]
|
||||
assert len(trainable) > 0, "layer 0 has no trainable parameters"
|
||||
|
||||
for i, p in enumerate(trainable):
|
||||
assert p.grad is not None, f"layer 0 param[{i}] has None grad"
|
||||
assert torch.isfinite(p.grad).all().item(), (
|
||||
f"layer 0 param[{i}] grad contains NaN/Inf")
|
||||
|
||||
any_nonzero = any(
|
||||
p.grad.detach().float().norm().item() > 0.0 for p in trainable)
|
||||
assert any_nonzero, (
|
||||
"all layer-0 grads are exactly zero; backward did not "
|
||||
"reach the first transformer block")
|
||||
|
||||
# Teacher forcing must build its own (concatenated) attention mask and
|
||||
# must not have constructed the diffusion-forcing mask.
|
||||
assert model.transformer.teacher_forcing_block_mask is not None, (
|
||||
"teacher-forcing mask was not constructed")
|
||||
assert model.transformer.block_mask is None, (
|
||||
"diffusion-forcing mask should not be built on the TF path")
|
||||
@@ -456,11 +456,16 @@ class ValidationCallback(Callback):
|
||||
None,
|
||||
)
|
||||
|
||||
loaded_modules: dict[str, Any] = {"transformer": transformer}
|
||||
# Distillation methods build the flow-match scheduler their few-step DMD
|
||||
# sampler needs; inject it so the pipeline doesn't fall back to UniPC.
|
||||
method_scheduler = getattr(self.method, "_sf_scheduler", None)
|
||||
if method_scheduler is not None:
|
||||
loaded_modules["scheduler"] = method_scheduler
|
||||
|
||||
kwargs: dict[str, Any] = {
|
||||
"inference_mode": True,
|
||||
"loaded_modules": {
|
||||
"transformer": transformer,
|
||||
},
|
||||
"loaded_modules": loaded_modules,
|
||||
"tp_size": tc.distributed.tp_size,
|
||||
"sp_size": tc.distributed.sp_size,
|
||||
"num_gpus": tc.distributed.num_gpus,
|
||||
@@ -477,12 +482,6 @@ class ValidationCallback(Callback):
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
scheduler = self._pipeline.get_module("scheduler")
|
||||
if (scheduler is not None and type(scheduler).__name__ == "SelfForcingFlowMatchScheduler"):
|
||||
scheduler.sigma_min = 0.0
|
||||
scheduler.extra_one_step = True
|
||||
scheduler.set_timesteps(num_inference_steps=1000, training=True)
|
||||
|
||||
self._pipeline_key = key
|
||||
return self._pipeline
|
||||
|
||||
|
||||
@@ -9,6 +9,8 @@ __all__ = [
|
||||
"KDMethod",
|
||||
"SelfForcingMethod",
|
||||
"DiffusionForcingSFTMethod",
|
||||
"TeacherForcingSFTMethod",
|
||||
"CausalConsistencyDistillationMethod",
|
||||
]
|
||||
|
||||
|
||||
@@ -28,4 +30,11 @@ def __getattr__(name: str) -> object:
|
||||
if name == "DiffusionForcingSFTMethod":
|
||||
from fastvideo.train.methods.fine_tuning.dfsft import DiffusionForcingSFTMethod
|
||||
return DiffusionForcingSFTMethod
|
||||
if name == "TeacherForcingSFTMethod":
|
||||
from fastvideo.train.methods.fine_tuning.tfsft import TeacherForcingSFTMethod
|
||||
return TeacherForcingSFTMethod
|
||||
if name == "CausalConsistencyDistillationMethod":
|
||||
from fastvideo.train.methods.consistency_model.causal_cd import (
|
||||
CausalConsistencyDistillationMethod, )
|
||||
return CausalConsistencyDistillationMethod
|
||||
raise AttributeError(name)
|
||||
|
||||
@@ -1,3 +1,22 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
__all__: list[str] = []
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from fastvideo.train.methods.consistency_model.causal_cd import (
|
||||
CausalConsistencyDistillationMethod, )
|
||||
|
||||
__all__ = [
|
||||
"CausalConsistencyDistillationMethod",
|
||||
]
|
||||
|
||||
|
||||
def __getattr__(name: str) -> object:
|
||||
if name == "CausalConsistencyDistillationMethod":
|
||||
from fastvideo.train.methods.consistency_model.causal_cd import (
|
||||
CausalConsistencyDistillationMethod, )
|
||||
|
||||
return CausalConsistencyDistillationMethod
|
||||
raise AttributeError(name)
|
||||
|
||||
@@ -0,0 +1,237 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Causal consistency distillation method (algorithm layer)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo.models.schedulers.scheduling_self_forcing_flow_match import (SelfForcingFlowMatchScheduler)
|
||||
from fastvideo.train.methods.base import LogScalar, TrainingMethod
|
||||
from fastvideo.train.models.base import ModelBase
|
||||
from fastvideo.train.utils.checkpoint import _FullModelState
|
||||
from fastvideo.train.utils.optimizer import build_optimizer_and_scheduler
|
||||
|
||||
|
||||
class CausalConsistencyDistillationMethod(TrainingMethod):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
cfg: Any,
|
||||
role_models: dict[str, ModelBase],
|
||||
) -> None:
|
||||
super().__init__(cfg=cfg, role_models=role_models)
|
||||
|
||||
for role in ("student", "teacher", "ema"):
|
||||
if role not in role_models:
|
||||
raise ValueError(f"Causal-CD requires role {role!r} "
|
||||
"(student trainable; teacher + ema frozen, "
|
||||
"both initialized from the student's "
|
||||
"checkpoint)")
|
||||
if not self.student._trainable:
|
||||
raise ValueError("Causal-CD requires student to be trainable")
|
||||
self.teacher = role_models["teacher"]
|
||||
self.ema_model = role_models["ema"]
|
||||
|
||||
self._attn_kind = self._infer_attn_kind()
|
||||
self._guidance_scale = float(self.method_config.get("guidance_scale", 3.0))
|
||||
self._discrete_cd_n = int(self.method_config.get("discrete_cd_N", 48))
|
||||
if self._discrete_cd_n < 2:
|
||||
raise ValueError("method.discrete_cd_N must be >= 2")
|
||||
self._ema_decay = float(self.method_config.get("ema_decay", 0.99))
|
||||
self._ema_start_step = int(self.method_config.get("ema_start_step", 200))
|
||||
shift = getattr(self.training_config.pipeline_config, "flow_shift", None)
|
||||
self._flow_shift = float(shift) if shift else 5.0
|
||||
|
||||
self.student.init_preprocessors(self.training_config)
|
||||
self._sf_scheduler = SelfForcingFlowMatchScheduler(
|
||||
num_inference_steps=self._discrete_cd_n,
|
||||
num_train_timesteps=int(self.student.num_train_timesteps),
|
||||
shift=self._flow_shift,
|
||||
sigma_min=0.0,
|
||||
sigma_max=1.0,
|
||||
extra_one_step=True,
|
||||
training=False,
|
||||
)
|
||||
self._init_optimizers_and_schedulers()
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@property
|
||||
def _optimizer_dict(self) -> dict[str, Any]:
|
||||
return {"student": self._student_optimizer}
|
||||
|
||||
@property
|
||||
def _lr_scheduler_dict(self) -> dict[str, Any]:
|
||||
return {"student": self._student_lr_scheduler}
|
||||
|
||||
def get_optimizers(self, iteration: int) -> list[torch.optim.Optimizer]:
|
||||
del iteration
|
||||
return [self._student_optimizer]
|
||||
|
||||
def get_lr_schedulers(self, iteration: int) -> list[Any]:
|
||||
del iteration
|
||||
return [self._student_lr_scheduler]
|
||||
|
||||
def checkpoint_state(self) -> dict[str, Any]:
|
||||
# The EMA role is frozen (so the base class skips it) but mutated by
|
||||
# _update_ema every step; without persisting it a resume reloads the
|
||||
# EMA from init_from and the consistency target snaps back to the
|
||||
# base checkpoint. Mirrors DiffusionNFT's frozen "old" role.
|
||||
states = super().checkpoint_state()
|
||||
states["roles.ema.transformer"] = _FullModelState(self.ema_model.transformer)
|
||||
return states
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def single_train_step(
|
||||
self,
|
||||
batch: dict[str, Any],
|
||||
iteration: int,
|
||||
) -> tuple[dict[str, torch.Tensor], dict[str, Any], dict[str, LogScalar]]:
|
||||
del iteration
|
||||
training_batch = self.student.prepare_batch(
|
||||
batch,
|
||||
generator=self.cuda_generator,
|
||||
latents_source="data",
|
||||
)
|
||||
clean_latents = training_batch.latents
|
||||
if not torch.is_tensor(clean_latents) or clean_latents.ndim != 5:
|
||||
raise ValueError("Causal-CD expects [B, T, C, H, W] latents")
|
||||
|
||||
batch_size, num_latents = int(clean_latents.shape[0]), int(clean_latents.shape[1])
|
||||
device = clean_latents.device
|
||||
|
||||
sigmas = self._sf_scheduler.sigmas.to(device)
|
||||
timesteps = self._sf_scheduler.timesteps.to(device)
|
||||
idx = torch.randint(0, self._discrete_cd_n - 1, (1, ), generator=self.cuda_generator, device=device).squeeze(0)
|
||||
t, t_next = timesteps[idx], timesteps[idx + 1]
|
||||
sigma_t, sigma_t_next = sigmas[idx], sigmas[idx + 1]
|
||||
t_pf = t * torch.ones(batch_size, num_latents, device=device)
|
||||
t_next_pf = t_next * torch.ones(batch_size, num_latents, device=device)
|
||||
|
||||
noise = torch.randn(
|
||||
clean_latents.shape,
|
||||
generator=self.cuda_generator,
|
||||
device=device,
|
||||
dtype=clean_latents.dtype,
|
||||
)
|
||||
latent_t = (1.0 - sigma_t) * clean_latents + sigma_t * noise
|
||||
|
||||
# Set before any forward: predict_noise feeds batch.timesteps into
|
||||
# set_forward_context (VSA sparsity gating), so the teacher CFG
|
||||
# passes below must not see the stale timesteps from prepare_batch.
|
||||
training_batch.timesteps = t_pf
|
||||
|
||||
with torch.no_grad():
|
||||
v_cond = self._predict_flow(self.teacher,
|
||||
latent_t,
|
||||
t_pf,
|
||||
training_batch,
|
||||
conditional=True,
|
||||
clean_x=clean_latents)
|
||||
v_uncond = self._predict_flow(self.teacher,
|
||||
latent_t,
|
||||
t_pf,
|
||||
training_batch,
|
||||
conditional=False,
|
||||
clean_x=clean_latents)
|
||||
v_pred = v_uncond + self._guidance_scale * (v_cond - v_uncond)
|
||||
dt = ((t - t_next) / float(self.student.num_train_timesteps))
|
||||
latent_t_next = latent_t - dt * v_pred
|
||||
|
||||
flow_student = self._predict_flow(self.student,
|
||||
latent_t,
|
||||
t_pf,
|
||||
training_batch,
|
||||
conditional=True,
|
||||
clean_x=clean_latents)
|
||||
x0_t = latent_t - sigma_t * flow_student
|
||||
|
||||
with torch.no_grad():
|
||||
flow_ema = self._predict_flow(self.ema_model,
|
||||
latent_t_next,
|
||||
t_next_pf,
|
||||
training_batch,
|
||||
conditional=True,
|
||||
clean_x=clean_latents)
|
||||
x0_t_next = latent_t_next - sigma_t_next * flow_ema
|
||||
|
||||
loss = F.mse_loss(x0_t.float(), x0_t_next.float())
|
||||
|
||||
loss_map = {"total_loss": loss, "causal_cd_loss": loss}
|
||||
attn_metadata = (training_batch.attn_metadata_vsa if self._attn_kind == "vsa" else training_batch.attn_metadata)
|
||||
outputs: dict[str, Any] = {"_fv_backward": (t_pf, attn_metadata)}
|
||||
metrics: dict[str, LogScalar] = {}
|
||||
return loss_map, outputs, metrics
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def backward(
|
||||
self,
|
||||
loss_map: dict[str, torch.Tensor],
|
||||
outputs: dict[str, Any],
|
||||
*,
|
||||
grad_accum_rounds: int = 1,
|
||||
) -> None:
|
||||
grad_accum_rounds = max(1, int(grad_accum_rounds))
|
||||
ctx = outputs.get("_fv_backward")
|
||||
if ctx is None:
|
||||
super().backward(loss_map, outputs, grad_accum_rounds=grad_accum_rounds)
|
||||
return
|
||||
self.student.backward(loss_map["total_loss"], ctx, grad_accum_rounds=grad_accum_rounds)
|
||||
|
||||
def optimizers_schedulers_step(self, iteration: int) -> None:
|
||||
super().optimizers_schedulers_step(iteration)
|
||||
if iteration >= self._ema_start_step:
|
||||
self._update_ema()
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Internal helpers
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _predict_flow(
|
||||
self,
|
||||
model: ModelBase,
|
||||
latents: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
batch: Any,
|
||||
*,
|
||||
conditional: bool,
|
||||
clean_x: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
return model.predict_noise(latents,
|
||||
timestep,
|
||||
batch,
|
||||
conditional=conditional,
|
||||
cfg_uncond=None,
|
||||
attn_kind=self._attn_kind,
|
||||
clean_x=clean_x)
|
||||
|
||||
@torch.no_grad()
|
||||
def _update_ema(self) -> None:
|
||||
decay = self._ema_decay
|
||||
for ema_p, p in zip(self.ema_model.transformer.parameters(), self.student.transformer.parameters(),
|
||||
strict=True):
|
||||
ema_p.mul_(decay).add_(p.detach().to(ema_p.dtype), alpha=1.0 - decay)
|
||||
|
||||
def _init_optimizers_and_schedulers(self) -> None:
|
||||
tc = self.training_config
|
||||
student_lr = float(tc.optimizer.learning_rate)
|
||||
if student_lr <= 0.0:
|
||||
raise ValueError("training.learning_rate must be > 0 for causal-cd")
|
||||
student_params = [p for p in self.student.transformer.parameters() if p.requires_grad]
|
||||
(
|
||||
self._student_optimizer,
|
||||
self._student_lr_scheduler,
|
||||
) = build_optimizer_and_scheduler(
|
||||
params=student_params,
|
||||
optimizer_config=tc.optimizer,
|
||||
loop_config=tc.loop,
|
||||
learning_rate=student_lr,
|
||||
betas=tc.optimizer.betas,
|
||||
scheduler_name=str(tc.optimizer.lr_scheduler),
|
||||
)
|
||||
@@ -7,10 +7,12 @@ from typing import TYPE_CHECKING
|
||||
if TYPE_CHECKING:
|
||||
from fastvideo.train.methods.fine_tuning.dfsft import DiffusionForcingSFTMethod
|
||||
from fastvideo.train.methods.fine_tuning.finetune import FineTuneMethod
|
||||
from fastvideo.train.methods.fine_tuning.tfsft import TeacherForcingSFTMethod
|
||||
|
||||
__all__ = [
|
||||
"DiffusionForcingSFTMethod",
|
||||
"FineTuneMethod",
|
||||
"TeacherForcingSFTMethod",
|
||||
]
|
||||
|
||||
|
||||
@@ -25,4 +27,9 @@ def __getattr__(name: str) -> object:
|
||||
from fastvideo.train.methods.fine_tuning.finetune import FineTuneMethod
|
||||
|
||||
return FineTuneMethod
|
||||
if name == "TeacherForcingSFTMethod":
|
||||
from fastvideo.train.methods.fine_tuning.tfsft import (
|
||||
TeacherForcingSFTMethod, )
|
||||
|
||||
return TeacherForcingSFTMethod
|
||||
raise AttributeError(name)
|
||||
|
||||
@@ -135,12 +135,11 @@ class DiffusionForcingSFTMethod(TrainingMethod):
|
||||
t_inhom.flatten(),
|
||||
)
|
||||
|
||||
pred = self.student.predict_noise(
|
||||
pred = self._predict_noise(
|
||||
noisy_latents,
|
||||
t_inhom,
|
||||
training_batch,
|
||||
conditional=True,
|
||||
attn_kind=self._attn_kind,
|
||||
clean_latents,
|
||||
)
|
||||
|
||||
if bool(self.training_config.model.precondition_outputs):
|
||||
@@ -178,6 +177,23 @@ class DiffusionForcingSFTMethod(TrainingMethod):
|
||||
metrics: dict[str, LogScalar] = {}
|
||||
return loss_map, outputs, metrics
|
||||
|
||||
def _predict_noise(
|
||||
self,
|
||||
noisy_latents: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
training_batch: Any,
|
||||
clean_latents: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
# Unused here; the teacher-forcing subclass overrides this to pass clean_x.
|
||||
del clean_latents
|
||||
return self.student.predict_noise(
|
||||
noisy_latents,
|
||||
timestep,
|
||||
training_batch,
|
||||
conditional=True,
|
||||
attn_kind=self._attn_kind,
|
||||
)
|
||||
|
||||
# TrainingMethod override: backward
|
||||
def backward(
|
||||
self,
|
||||
|
||||
@@ -0,0 +1,30 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Teacher-forcing SFT method (TFSFT; algorithm layer)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.train.methods.fine_tuning.dfsft import (
|
||||
DiffusionForcingSFTMethod, )
|
||||
|
||||
|
||||
class TeacherForcingSFTMethod(DiffusionForcingSFTMethod):
|
||||
|
||||
def _predict_noise(
|
||||
self,
|
||||
noisy_latents: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
training_batch: Any,
|
||||
clean_latents: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
return self.student.predict_noise(
|
||||
noisy_latents,
|
||||
timestep,
|
||||
training_batch,
|
||||
conditional=True,
|
||||
attn_kind=self._attn_kind,
|
||||
clean_x=clean_latents,
|
||||
)
|
||||
@@ -10,12 +10,6 @@ from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from torch.distributed.checkpoint.state_dict import (
|
||||
StateDictOptions,
|
||||
get_model_state_dict,
|
||||
set_model_state_dict,
|
||||
)
|
||||
from torch.distributed.checkpoint.stateful import Stateful
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
from fastvideo.dataset.parquet_dataset_map_style import (
|
||||
@@ -37,6 +31,7 @@ from fastvideo.train.methods.rl.common import (
|
||||
validation_caption,
|
||||
validation_shard_indices,
|
||||
)
|
||||
from fastvideo.train.utils.checkpoint import _FullModelState
|
||||
from fastvideo.train.utils.config import (
|
||||
get_optional_float,
|
||||
get_optional_int,
|
||||
@@ -78,31 +73,6 @@ class _DiffusionNFTEMAState:
|
||||
self._method._ema_update_count = int(update_count)
|
||||
|
||||
|
||||
class _FullModelState(Stateful):
|
||||
"""DCP wrapper that saves frozen model parameters too.
|
||||
|
||||
The shared ``ModelWrapper`` intentionally filters to ``requires_grad``
|
||||
parameters. DiffusionNFT's old policy is frozen but must be restored on
|
||||
resume, so it needs full model state.
|
||||
"""
|
||||
|
||||
def __init__(self, model: torch.nn.Module) -> None:
|
||||
self.model = model
|
||||
|
||||
def state_dict(self) -> dict[str, Any]:
|
||||
return get_model_state_dict(self.model) # type: ignore[no-any-return]
|
||||
|
||||
def load_state_dict(
|
||||
self,
|
||||
state_dict: dict[str, Any],
|
||||
) -> None:
|
||||
set_model_state_dict(
|
||||
self.model,
|
||||
model_state_dict=state_dict,
|
||||
options=StateDictOptions(strict=False),
|
||||
)
|
||||
|
||||
|
||||
class DiffusionNFTMethod(TrainingMethod):
|
||||
"""DiffusionNFT-style RL for diffusion models.
|
||||
|
||||
|
||||
@@ -320,6 +320,8 @@ class WanModel(ModelBase):
|
||||
conditional: bool,
|
||||
cfg_uncond: dict[str, Any] | None = None,
|
||||
attn_kind: Literal["dense", "vsa"] = "dense",
|
||||
clean_x: torch.Tensor | None = None,
|
||||
aug_t: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
device_type = self.device.type
|
||||
dtype = self._get_training_dtype()
|
||||
@@ -347,7 +349,11 @@ class WanModel(ModelBase):
|
||||
current_timestep=batch.timesteps,
|
||||
attn_metadata=attn_metadata,
|
||||
):
|
||||
input_kwargs = (self._build_distill_input_kwargs(noisy_latents, timestep, text_dict))
|
||||
input_kwargs = (self._build_distill_input_kwargs(noisy_latents,
|
||||
timestep,
|
||||
text_dict,
|
||||
clean_x=clean_x,
|
||||
aug_t=aug_t))
|
||||
transformer = self._get_transformer(timestep)
|
||||
pred_noise = transformer(**input_kwargs).permute(0, 2, 1, 3, 4)
|
||||
return pred_noise
|
||||
@@ -530,17 +536,24 @@ class WanModel(ModelBase):
|
||||
noise_input: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
text_dict: dict[str, torch.Tensor] | None,
|
||||
clean_x: torch.Tensor | None = None,
|
||||
aug_t: torch.Tensor | None = None,
|
||||
) -> dict[str, Any]:
|
||||
if text_dict is None:
|
||||
raise ValueError("text_dict cannot be None for "
|
||||
"Wan distillation")
|
||||
return {
|
||||
kwargs: dict[str, Any] = {
|
||||
"hidden_states": noise_input.permute(0, 2, 1, 3, 4),
|
||||
"encoder_hidden_states": text_dict["encoder_hidden_states"],
|
||||
"encoder_attention_mask": text_dict["encoder_attention_mask"],
|
||||
"timestep": timestep,
|
||||
"return_dict": False,
|
||||
}
|
||||
if clean_x is not None:
|
||||
# Teacher forcing: clean context latents (+ optional aug timestep).
|
||||
kwargs["clean_x"] = clean_x.permute(0, 2, 1, 3, 4)
|
||||
kwargs["aug_t"] = aug_t
|
||||
return kwargs
|
||||
|
||||
def _get_transformer(self, timestep: torch.Tensor) -> torch.nn.Module:
|
||||
return self.transformer
|
||||
|
||||
@@ -49,6 +49,7 @@ class WanCausalModel(WanModel, CausalModelBase):
|
||||
transformer_override_safetensor: str
|
||||
| None = None,
|
||||
lora: LoraConfig | dict[str, Any] | None = None,
|
||||
num_frames_per_block: int | None = None,
|
||||
) -> None:
|
||||
super().__init__(
|
||||
init_from=init_from,
|
||||
@@ -62,6 +63,16 @@ class WanCausalModel(WanModel, CausalModelBase):
|
||||
)
|
||||
self._streaming_caches: (dict[tuple[int, str], _StreamingCaches]) = {}
|
||||
|
||||
if num_frames_per_block is not None:
|
||||
num_frames_per_block = int(num_frames_per_block)
|
||||
if not 1 <= num_frames_per_block <= 3:
|
||||
# Same bound as CausalWanTransformer3DModel's config path
|
||||
# (assert num_frame_per_block <= 3); this override must not
|
||||
# bypass it.
|
||||
raise ValueError("num_frames_per_block must be between 1 and 3, "
|
||||
f"got {num_frames_per_block}")
|
||||
self.transformer.num_frame_per_block = num_frames_per_block
|
||||
|
||||
# --- CausalModelBase override: clear_caches ---
|
||||
def clear_caches(
|
||||
self,
|
||||
|
||||
@@ -15,6 +15,13 @@ import numpy as np
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.distributed.checkpoint as dcp
|
||||
from torch.distributed.checkpoint.state_dict import (
|
||||
StateDictOptions,
|
||||
get_model_state_dict,
|
||||
set_model_state_dict,
|
||||
)
|
||||
from torch.distributed.checkpoint.stateful import Stateful
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -131,6 +138,32 @@ class _RoleModuleContainer(torch.nn.Module):
|
||||
self.add_module(name, module)
|
||||
|
||||
|
||||
class _FullModelState(Stateful):
|
||||
"""DCP wrapper that saves frozen model parameters too.
|
||||
|
||||
The shared ``ModelWrapper`` intentionally filters to ``requires_grad``
|
||||
parameters. Frozen-but-mutated roles (e.g. DiffusionNFT's old policy,
|
||||
causal-CD's EMA target) must still be restored on resume, so they need
|
||||
full model state.
|
||||
"""
|
||||
|
||||
def __init__(self, model: torch.nn.Module) -> None:
|
||||
self.model = model
|
||||
|
||||
def state_dict(self) -> dict[str, Any]:
|
||||
return get_model_state_dict(self.model) # type: ignore[no-any-return]
|
||||
|
||||
def load_state_dict(
|
||||
self,
|
||||
state_dict: dict[str, Any],
|
||||
) -> None:
|
||||
set_model_state_dict(
|
||||
self.model,
|
||||
model_state_dict=state_dict,
|
||||
options=StateDictOptions(strict=False),
|
||||
)
|
||||
|
||||
|
||||
class _CallbackStateWrapper:
|
||||
"""Wraps a CallbackDict for DCP save/load."""
|
||||
|
||||
|
||||
Reference in New Issue
Block a user