Compare commits

..
Author SHA1 Message Date
SolitaryThinker 80d5b16e1b fix sp i2v2 2025-06-29 20:24:38 -07:00
SolitaryThinker 7dcfe4ea8f fix sp i2v 2025-06-29 19:46:08 -07:00
SolitaryThinker 9ab7f031af fix sp i2v 2025-06-29 19:41:29 -07:00
19 changed files with 100 additions and 193 deletions
+20 -10
View File
@@ -1,12 +1,13 @@
env:
IMAGE_VERSION: "py3.12-latest"
BUILDKITE_CLEAN_CHECKOUT: true
steps:
- label: "pre-commit"
command: ".buildkite/scripts/pre_commit.sh"
agents:
queue: "default"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- wait
@@ -22,9 +23,10 @@ steps:
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "Encoder Tests"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=encoder
agents:
queue: "default"
@@ -35,9 +37,10 @@ steps:
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "VAE Tests"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=vae
agents:
queue: "default"
@@ -50,18 +53,20 @@ steps:
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "Transformer Tests"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=transformer
agents:
queue: "default"
- path:
- "fastvideo/v1/**/*.py"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
command: "timeout 60m .buildkite/scripts/pr_test.sh"
label: "SSIM Tests"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=ssim
agents:
queue: "default"
@@ -70,9 +75,10 @@ steps:
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "Training Tests"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=training
agents:
queue: "default"
@@ -86,9 +92,10 @@ steps:
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "Training Tests VSA"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=training_vsa
agents:
queue: "default"
@@ -101,9 +108,10 @@ steps:
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "Inference Tests STA"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=inference_sta
agents:
queue: "default"
@@ -115,9 +123,10 @@ steps:
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "Precision Tests STA"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=precision_sta
agents:
queue: "default"
@@ -130,9 +139,10 @@ steps:
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "Precision Tests VSA"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=precision_vsa
agents:
queue: "default"
+1 -1
View File
@@ -58,7 +58,7 @@ if [ -z "${TEST_TYPE:-}" ]; then
fi
log "Test type: $TEST_TYPE"
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT BUILDKITE_PULL_REQUEST=$BUILDKITE_PULL_REQUEST IMAGE_VERSION=$IMAGE_VERSION"
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT IMAGE_VERSION=$IMAGE_VERSION"
case "$TEST_TYPE" in
"encoder")
-14
View File
@@ -15,11 +15,6 @@ With FastVideo's optimizations, you can achieve more than 3x inference improveme
<img src=assets/perf.png width="90%"/>
</div>
## NEWS
- ```2025/06/14```: Release finetuning and inference code for [VSA](https://arxiv.org/pdf/2505.13389)
- ```2025/04/24```: [FastVideo V1](https://hao-ai-lab.github.io/blogs/fastvideo/) is released!
- ```2025/02/18```: Release the inference code for [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
## Key Features
FastVideo has the following features:
@@ -133,15 +128,6 @@ We thank MBZUAI and [Anyscale](https://www.anyscale.com/) for their support thro
If you use FastVideo for your research, please cite our paper:
```bibtex
@misc{zhang2025vsafastervideodiffusion,
title={VSA: Faster Video Diffusion with Trainable Sparse Attention},
author={Peiyuan Zhang and Haofeng Huang and Yongqi Chen and Will Lin and Zhengzhong Liu and Ion Stoica and Eric Xing and Hao Zhang},
year={2025},
eprint={2505.13389},
archivePrefix={arXiv},
primaryClass={cs.CV},
url={https://arxiv.org/abs/2505.13389},
}
@misc{zhang2025fastvideogenerationsliding,
title={Fast Video Generation with Sliding Tile Attention},
author={Peiyuan Zhang and Yongqi Chen and Runlong Su and Hangliang Ding and Ion Stoica and Zhenghong Liu and Hao Zhang},
@@ -22,5 +22,4 @@ torchrun --nproc_per_node=$GPU_NUM \
--validation_dataset_file $VALIDATION_PATH \
--samples_per_file 8 \
--flush_frequency 8 \
--video_length_tolerance_range 5 \
--preprocess_task "i2v"
@@ -22,5 +22,4 @@ torchrun --nproc_per_node=$GPU_NUM \
--validation_dataset_file $VALIDATION_PATH \
--samples_per_file 8 \
--flush_frequency 8 \
--video_length_tolerance_range 5 \
--preprocess_task "t2v"
+1 -1
View File
@@ -52,7 +52,7 @@ class PipelineConfig:
# VAE configuration
vae_config: VAEConfig = field(default_factory=VAEConfig)
vae_precision: str = "fp32"
vae_precision: str = "fp16"
vae_tiling: bool = True
vae_sp: bool = True
+1 -1
View File
@@ -50,7 +50,7 @@ class WanT2V480PConfig(PipelineConfig):
# Precision for each component
precision: str = "bf16"
vae_precision: str = "fp32"
vae_precision: str = "fp16"
text_encoder_precisions: Tuple[str, ...] = field(
default_factory=lambda: ("fp32", ))
@@ -4,7 +4,7 @@
"use_cpu_offload": true,
"disable_autocast": false,
"precision": "bf16",
"vae_precision": "fp32",
"vae_precision": "fp16",
"vae_tiling": false,
"vae_sp": false,
"vae_config": {
@@ -4,7 +4,7 @@
"use_cpu_offload": true,
"disable_autocast": false,
"precision": "bf16",
"vae_precision": "fp32",
"vae_precision": "fp16",
"vae_tiling": false,
"vae_sp": false,
"vae_config": {
+4 -6
View File
@@ -413,7 +413,7 @@ class TrainingArgs(FastVideoArgs):
lr_scheduler: str = "constant"
lr_warmup_steps: int = 0
max_grad_norm: float = 0.0
enable_gradient_checkpointing_type: Optional[str] = None
gradient_checkpointing: bool = False
selective_checkpointing: float = 0.0
allow_tf32: bool = False
mixed_precision: str = ""
@@ -612,11 +612,9 @@ class TrainingArgs(FastVideoArgs):
parser.add_argument("--max-grad-norm",
type=float,
help="Maximum gradient norm")
parser.add_argument("--enable-gradient-checkpointing-type",
type=str,
choices=["full", "ops", "block_skip"],
default=None,
help="Gradient checkpointing type")
parser.add_argument("--gradient-checkpointing",
action=StoreBoolean,
help="Whether to use gradient checkpointing")
parser.add_argument("--selective-checkpointing",
type=float,
help="Selective checkpointing threshold")
+6 -7
View File
@@ -71,16 +71,15 @@ class BaseLayerWithLoRA(nn.Module):
f"cuda:{torch.cuda.current_device()}").full_tensor()
data += (self.slice_lora_b_weights(self.lora_B)
@ self.slice_lora_a_weights(self.lora_A)).to(data)
self.base_layer.weight = nn.Parameter(
distribute_tensor(data, mesh,
placements=placements).to(current_device))
self.base_layer.weight.data = distribute_tensor(
data, mesh, placements=placements).to(current_device)
else:
current_device = self.base_layer.weight.data.device
data = self.base_layer.weight.to(
data = self.base_layer.weight.data.to(
f"cuda:{torch.cuda.current_device()}")
data += \
(self.slice_lora_b_weights(self.lora_B) @ self.slice_lora_a_weights(self.lora_A)).to(data)
self.base_layer.weight = nn.Parameter(data.to(current_device))
self.base_layer.weight.data = data.to(current_device)
self.merged = True
@torch.no_grad()
@@ -107,8 +106,8 @@ class BaseLayerWithLoRA(nn.Module):
f"cuda:{torch.cuda.current_device()}").full_tensor()
data -= self.slice_lora_b_weights(
self.lora_B) @ self.slice_lora_a_weights(self.lora_A)
self.base_layer.weight = nn.Parameter(
distribute_tensor(data, mesh, placements=placement).to(device))
self.base_layer.weight.data = distribute_tensor(
data, mesh, placements=placement).to(device)
else:
self.base_layer.weight.data -= \
self.slice_lora_b_weights(self.lora_B) @\
+11 -24
View File
@@ -18,8 +18,6 @@ image = (
"PATH": "/root/.cargo/bin:$PATH",
"BUILDKITE_REPO": os.environ.get("BUILDKITE_REPO", ""),
"BUILDKITE_COMMIT": os.environ.get("BUILDKITE_COMMIT", ""),
"BUILDKITE_PULL_REQUEST": os.environ.get("BUILDKITE_PULL_REQUEST", ""),
"IMAGE_VERSION": os.environ.get("IMAGE_VERSION", ""),
})
)
@@ -31,27 +29,16 @@ def run_test(pytest_command: str):
git_repo = os.environ.get("BUILDKITE_REPO")
git_commit = os.environ.get("BUILDKITE_COMMIT")
pr_number = os.environ.get("BUILDKITE_PULL_REQUEST")
print(f"Cloning repository: {git_repo}")
print(f"Target commit: {git_commit}")
if pr_number:
print(f"PR number: {pr_number}")
# For PRs (including forks), use GitHub's PR refs to get the correct commit
if pr_number and pr_number != "false":
checkout_command = f"git fetch --prune origin refs/pull/{pr_number}/head && git checkout FETCH_HEAD"
print(f"Using PR ref for checkout: {checkout_command}")
else:
checkout_command = f"git checkout {git_commit}"
print(f"Using direct commit checkout: {checkout_command}")
print(f"Checking out commit: {git_commit}")
command = f"""
source $HOME/.local/bin/env &&
source /opt/venv/bin/activate &&
git clone {git_repo} /FastVideo &&
cd /FastVideo &&
{checkout_command} &&
git checkout {git_commit} &&
uv pip install -e .[test] &&
{pytest_command}
"""
@@ -62,38 +49,38 @@ def run_test(pytest_command: str):
sys.exit(result.returncode)
@app.function(gpu="L40S:1", image=image, timeout=900)
@app.function(gpu="L40S:1", image=image, timeout=1800)
def run_encoder_tests():
run_test("pytest ./fastvideo/v1/tests/encoders -vs")
@app.function(gpu="L40S:1", image=image, timeout=900)
@app.function(gpu="L40S:1", image=image, timeout=1800)
def run_vae_tests():
run_test("pytest ./fastvideo/v1/tests/vaes -vs")
@app.function(gpu="L40S:1", image=image, timeout=900)
@app.function(gpu="L40S:1", image=image, timeout=1800)
def run_transformer_tests():
run_test("pytest ./fastvideo/v1/tests/transformers -vs")
@app.function(gpu="L40S:2", image=image, timeout=1800)
@app.function(gpu="L40S:2", image=image, timeout=3600)
def run_ssim_tests():
run_test("pytest ./fastvideo/v1/tests/ssim -vs")
@app.function(gpu="L40S:4", image=image, timeout=900, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
@app.function(gpu="L40S:4", image=image, timeout=1800, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
def run_training_tests():
run_test("wandb login $WANDB_API_KEY && pytest ./fastvideo/v1/tests/training/Vanilla -srP")
@app.function(gpu="H100:2", image=image, timeout=900, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
@app.function(gpu="H100:2", image=image, timeout=1800, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
def run_training_tests_VSA():
run_test("wandb login $WANDB_API_KEY && pytest ./fastvideo/v1/tests/training/VSA -srP")
@app.function(gpu="H100:2", image=image, timeout=900)
@app.function(gpu="H100:2", image=image, timeout=1800)
def run_inference_tests_STA():
run_test("pytest ./fastvideo/v1/tests/inference/STA -srP")
@app.function(gpu="H100:1", image=image, timeout=900)
@app.function(gpu="H100:1", image=image, timeout=1800)
def run_precision_tests_STA():
run_test("python csrc/attn/tests/test_sta.py")
@app.function(gpu="H100:1", image=image, timeout=900)
@app.function(gpu="H100:1", image=image, timeout=1800)
def run_precision_tests_VSA():
run_test("python csrc/attn/tests/test_block_sparse.py")
@@ -1,91 +0,0 @@
import collections
from enum import Enum
from typing import Optional
import torch
from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import (
checkpoint_wrapper)
TRANSFORMER_BLOCK_NAMES = [
"blocks",
"double_blocks",
"single_blocks",
"transformer_blocks",
"temporal_transformer_blocks",
"transformer_double_blocks",
"transformer_single_blocks",
]
class CheckpointType(str, Enum):
FULL = "full"
OPS = "ops"
BLOCK_SKIP = "block_skip"
_SELECTIVE_ACTIVATION_CHECKPOINTING_OPS = {
torch.ops.aten.mm.default,
torch.ops.aten._scaled_dot_product_efficient_attention.default,
torch.ops.aten._scaled_dot_product_flash_attention.default,
torch.ops._c10d_functional.reduce_scatter_tensor.default,
}
def apply_activation_checkpointing(
module: torch.nn.Module,
checkpointing_type: str = CheckpointType.FULL,
n_layer: int = 1) -> torch.nn.Module:
if checkpointing_type == CheckpointType.FULL:
module = _apply_activation_checkpointing_blocks(module)
elif checkpointing_type == CheckpointType.OPS:
module = _apply_activation_checkpointing_ops(
module, _SELECTIVE_ACTIVATION_CHECKPOINTING_OPS)
elif checkpointing_type == CheckpointType.BLOCK_SKIP:
module = _apply_activation_checkpointing_blocks(module, n_layer)
else:
raise ValueError(
f"Checkpointing type '{checkpointing_type}' not supported. Supported types are {CheckpointType.__members__.keys()}"
)
return module
def _apply_activation_checkpointing_blocks(
module: torch.nn.Module,
n_layer: Optional[int] = None) -> torch.nn.Module:
for transformer_block_name in TRANSFORMER_BLOCK_NAMES:
blocks: torch.nn.Module = getattr(module, transformer_block_name, None)
if blocks is None:
continue
for index, (layer_id, block) in enumerate(blocks.named_children()):
if n_layer is None or index % n_layer == 0:
block = checkpoint_wrapper(block, preserve_rng_state=False)
blocks.register_module(layer_id, block)
return module
def _apply_activation_checkpointing_ops(module: torch.nn.Module,
ops) -> torch.nn.Module:
from torch.utils.checkpoint import (CheckpointPolicy,
create_selective_checkpoint_contexts)
def _get_custom_policy(meta: dict[str, int]) -> CheckpointPolicy:
def _custom_policy(ctx, func, *args, **kwargs):
mode = "recompute" if ctx.is_recompute else "forward"
mm_count_key = f"{mode}_mm_count"
if func == torch.ops.aten.mm.default:
meta[mm_count_key] += 1
# Saves output of all compute ops, except every second mm
to_save = func in ops and not (func == torch.ops.aten.mm.default
and meta[mm_count_key] % 2 == 0)
return CheckpointPolicy.MUST_SAVE if to_save else CheckpointPolicy.PREFER_RECOMPUTE
return _custom_policy
def selective_checkpointing_context_fn():
meta: dict[str, int] = collections.defaultdict(int)
return create_selective_checkpoint_contexts(_get_custom_policy(meta))
return checkpoint_wrapper(module,
context_fn=selective_checkpointing_context_fn,
preserve_rng_state=False)
+11 -23
View File
@@ -31,8 +31,6 @@ from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines import (ComposedPipelineBase, ForwardBatch,
TrainingBatch)
from fastvideo.v1.training.activation_checkpoint import (
apply_activation_checkpointing)
from fastvideo.v1.training.training_utils import (
clip_grad_norm_while_handling_failing_dtensor_cases,
compute_density_for_timestep_sampling, get_sigmas, load_checkpoint,
@@ -84,11 +82,6 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
self.transformer.requires_grad_(True)
self.transformer.train()
if training_args.enable_gradient_checkpointing_type is not None:
self.transformer = apply_activation_checkpointing(
self.transformer,
checkpointing_type=training_args.
enable_gradient_checkpointing_type)
noise_scheduler = self.modules["scheduler"]
params_to_optimize = self.transformer.parameters()
@@ -316,18 +309,17 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
current_timestep=training_batch.current_timestep,
attn_metadata=training_batch.attn_metadata):
model_pred = self.transformer(**input_kwargs)
if self.training_args.precondition_outputs:
model_pred = training_batch.noisy_model_input - model_pred * training_batch.sigmas
target = training_batch.latents if self.training_args.precondition_outputs else training_batch.noise - training_batch.latents
if self.training_args.precondition_outputs:
model_pred = training_batch.noisy_model_input - model_pred * training_batch.sigmas
target = training_batch.latents if self.training_args.precondition_outputs else training_batch.noise - training_batch.latents
# make sure no implicit broadcasting happens
assert model_pred.shape == target.shape, f"model_pred.shape: {model_pred.shape}, target.shape: {target.shape}"
loss = (torch.mean((model_pred.float() - target.float())**2) /
self.training_args.gradient_accumulation_steps)
loss.backward()
avg_loss = loss.detach().clone()
# make sure no implicit broadcasting happens
assert model_pred.shape == target.shape, f"model_pred.shape: {model_pred.shape}, target.shape: {target.shape}"
loss = (torch.mean((model_pred.float() - target.float())**2) /
self.training_args.gradient_accumulation_steps)
loss.backward()
avg_loss = loss.detach().clone()
# logger.info(f"rank: {self.rank}, avg_loss: {avg_loss.item()}",
# local_main_process_only=False)
world_group = get_world_group()
@@ -554,8 +546,6 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
negative_prompt_attention_mask: torch.Tensor | None
) -> ForwardBatch:
assert len(validation_batch['info_list']
) == 1, "Only batch size 1 is supported for validation"
prompt = validation_batch['info_list'][0]['prompt']
prompt_embeds = validation_batch['text_embedding']
prompt_attention_mask = validation_batch['text_attention_mask']
@@ -639,7 +629,7 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
# Process each validation prompt for each validation step
for num_inference_steps in validation_steps:
step_videos: List[np.ndarray] = []
step_captions: List[str] = []
step_captions: List[str | None] = []
for validation_batch in validation_dataloader:
batch = self._prepare_validation_inputs(
@@ -647,9 +637,7 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
num_inference_steps, negative_prompt_embeds,
negative_prompt_attention_mask)
assert batch.prompt is not None and isinstance(
batch.prompt, str)
step_captions.append(batch.prompt)
step_captions.extend([None]) # TODO(peiyuan): add caption
# Run validation inference
with torch.no_grad(), torch.autocast("cuda",
+5 -5
View File
@@ -162,8 +162,8 @@ def save_checkpoint(transformer,
weight_path,
local_main_process_only=False)
# Convert fastvideo custom format to diffusers format and save
diffusers_state_dict = convert_custom_format_to_diffusers_format(
# Convert training format to diffusers format and save
diffusers_state_dict = convert_training_to_diffusers_format(
cpu_state, transformer)
save_file(diffusers_state_dict, weight_path)
@@ -487,10 +487,10 @@ def _has_foreach_support(tensors: List[torch.Tensor],
t is None or type(t) in [torch.Tensor] for t in tensors)
def convert_custom_format_to_diffusers_format(state_dict: Dict[str, Any],
transformer) -> Dict[str, Any]:
def convert_training_to_diffusers_format(state_dict: Dict[str, Any],
transformer) -> Dict[str, Any]:
"""
Convert fastvideo custom format state dict to diffusers format using reverse_param_names_mapping.
Convert training format state dict to diffusers format using reverse_param_names_mapping.
Args:
state_dict: State dict in training format
@@ -4,11 +4,12 @@ from copy import deepcopy
from typing import Any, Dict
import torch
import torch.distributed
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.dataset.dataloader.schema import (
pyarrow_schema_i2v, pyarrow_schema_i2v_validation)
from fastvideo.v1.distributed import get_torch_device
from fastvideo.v1.distributed import get_torch_device, get_sp_group
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.schedulers.scheduling_flow_unipc_multistep import (
@@ -18,6 +19,8 @@ from fastvideo.v1.pipelines.pipeline_batch_info import (ForwardBatch,
from fastvideo.v1.pipelines.wan.wan_i2v_pipeline import (
WanImageToVideoValidationPipeline)
from fastvideo.v1.training.training_pipeline import TrainingPipeline
from fastvideo.v1.training.training_utils import (shard_latents_across_sp,
clip_grad_norm_while_handling_failing_dtensor_cases)
from fastvideo.v1.utils import is_vsa_available
vsa_available = is_vsa_available()
@@ -204,6 +207,36 @@ class WanI2VTrainingPipeline(TrainingPipeline):
)
return batch
def _clip_grad_norm(self, training_batch: TrainingBatch) -> TrainingBatch:
"""Override to add gradient synchronization across SP ranks."""
assert self.training_args is not None
max_grad_norm = self.training_args.max_grad_norm
# CRITICAL FIX: Synchronize gradients across SP ranks before clipping
# Different SP ranks compute different gradients due to different noise patterns
# These gradients must be averaged across SP ranks for stable training
if self.training_args.sp_size > 1:
sp_group = get_sp_group()
for param in self.transformer.parameters():
if param.grad is not None:
# Average gradients across SP ranks
sp_group.all_reduce(param.grad, op=torch.distributed.ReduceOp.AVG)
if max_grad_norm is not None:
model_parts = [self.transformer]
grad_norm = clip_grad_norm_while_handling_failing_dtensor_cases(
[p for m in model_parts for p in m.parameters()],
max_grad_norm,
foreach=None,
)
assert grad_norm is not float('nan') or grad_norm is not float(
'inf')
grad_norm = grad_norm.item() if grad_norm is not None else 0.0
else:
grad_norm = 0.0
training_batch.grad_norm = grad_norm
return training_batch
def main(args) -> None:
logger.info("Starting training pipeline...")
+2 -1
View File
@@ -1,10 +1,11 @@
# trigger test
[build-system]
requires = ["setuptools>=61.0"]
build-backend = "setuptools.build_meta"
[project]
name = "fastvideo"
version = "0.1.1"
version = "0.1.0"
description = "FastVideo"
readme = "README.md"
requires-python = ">=3.8"
+1 -2
View File
@@ -48,5 +48,4 @@ torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
--weight_decay 0.01 \
--not_apply_cfg_solver \
--dit_precision "fp32" \
--max_grad_norm 1.0 \
--enable_gradient_checkpointing_type "full"
--max_grad_norm 1.0
+1 -2
View File
@@ -58,6 +58,5 @@ torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
--VSA_decay_sparsity 0.9 \
--VSA_decay_rate 0.03 \
--VSA_decay_interval_steps 30 \
--VSA_val_sparsity 0.9 \
--enable_gradient_checkpointing_type "full"
--VSA_val_sparsity 0.9
# --resume_from_checkpoint "$CHECKPOINT_PATH"