Compare commits

...
3 changed files with 80 additions and 13 deletions
@@ -0,0 +1,58 @@
# SPDX-License-Identifier: Apache-2.0
"""Guard the Modal FA4 defaults that keep CI lanes on their intended backend.
Pure text/AST analysis: no fastvideo imports, no torch, no Modal client.
"""
from __future__ import annotations
import ast
from pathlib import Path
REPO_ROOT = Path(__file__).resolve().parents[3]
MODAL_ROOT = REPO_ROOT / "fastvideo" / "tests" / "modal"
PR_TEST = MODAL_ROOT / "pr_test.py"
LAUNCH_L40S_JOB = MODAL_ROOT / "launch_l40s_job.py"
SSIM_TEST = MODAL_ROOT / "ssim_test.py"
def _function_strings(path: Path, function_name: str) -> str:
source = path.read_text(encoding="utf-8")
tree = ast.parse(source)
for node in tree.body:
if isinstance(node, ast.FunctionDef) and node.name == function_name:
return "\n".join(
child.value
for child in ast.walk(node)
if isinstance(child, ast.Constant)
and isinstance(child.value, str)
)
raise AssertionError(f"{function_name} not found in {path}")
def test_generic_l40s_launcher_defaults_fa4_off():
source = LAUNCH_L40S_JOB.read_text(encoding="utf-8")
assert '"FASTVIDEO_FA4": os.environ.get("FASTVIDEO_FA4", "0")' in source
def test_ssim_launcher_keeps_fa4_enabled_by_default():
source = SSIM_TEST.read_text(encoding="utf-8")
assert '"FASTVIDEO_FA4": os.environ.get("FASTVIDEO_FA4", "1")' in source
def test_pr_model_load_and_training_lanes_disable_fa4():
lanes = {
"run_transformer_tests": "pytest ./fastvideo/tests/transformers -vs",
"run_training_tests": "pytest ./fastvideo/tests/training/Vanilla -srP",
"run_training_lora_tests": "pytest ./fastvideo/tests/training/lora/test_lora_training.py -srP",
"run_training_tests_VSA": "pytest ./fastvideo/tests/training/VSA -srP",
"run_distill_dmd_tests": "pytest ./fastvideo/tests/training/distill/test_distill_dmd.py -vs",
"run_self_forcing_tests": "pytest ./fastvideo/tests/training/self-forcing/test_self_forcing.py -vs",
"run_train_framework_tests": "pytest ./fastvideo/tests/train/models ./fastvideo/tests/train/methods -vs",
"seed_grad_norm_references": "pytest ./fastvideo/tests/train/methods -vs -rs",
}
for function_name, pytest_command in lanes.items():
function_strings = _function_strings(PR_TEST, function_name)
assert "FASTVIDEO_FA4=0" in function_strings
assert pytest_command in function_strings
+4 -3
View File
@@ -97,9 +97,10 @@ image = (
"TOKENIZERS_PARALLELISM": "false",
**({"UV_TORCH_BACKEND": uv_torch_backend_override} if uv_torch_backend_override else {}),
"FASTVIDEO_ATTENTION_BACKEND": os.environ.get("FASTVIDEO_ATTENTION_BACKEND", "FLASH_ATTN"),
# FA4 is opt-in (FASTVIDEO_FA4); keep CI parity with the seeded
# references. Caller override wins.
"FASTVIDEO_FA4": os.environ.get("FASTVIDEO_FA4", "1"),
# FA4 is opt-in (FASTVIDEO_FA4). Generic ad hoc jobs should follow the
# product default unless a caller opts in through the local env or
# --env-vars.
"FASTVIDEO_FA4": os.environ.get("FASTVIDEO_FA4", "0"),
})
)
+18 -10
View File
@@ -73,8 +73,10 @@ ci_env_secret = modal.Secret.from_dict({
**({
"UV_TORCH_BACKEND": uv_torch_backend_override
} if uv_torch_backend_override else {}),
# FA4 is opt-in (FASTVIDEO_FA4); CI lanes keep it enabled to match the
# SSIM/perf baselines. Caller override wins.
# FA4 is opt-in (FASTVIDEO_FA4). Keep the default enabled for
# inference/perf parity; model-load and training lanes that do not exercise
# FA4 explicitly set FASTVIDEO_FA4=0 in their command strings below.
# Caller override wins.
"FASTVIDEO_FA4": os.environ.get("FASTVIDEO_FA4", "1"),
})
@@ -184,7 +186,8 @@ def run_vae_tests():
volumes={"/root/data": model_vol})
def run_transformer_tests():
run_test(
"export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/transformers -vs"
"export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && "
"FASTVIDEO_FA4=0 pytest ./fastvideo/tests/transformers -vs"
)
@@ -197,7 +200,8 @@ def run_transformer_tests():
volumes={"/root/data": model_vol})
def run_training_tests():
run_test(
"export HF_HOME='/root/data/.cache' && wandb login $WANDB_API_KEY && pytest ./fastvideo/tests/training/Vanilla -srP"
"export HF_HOME='/root/data/.cache' && wandb login $WANDB_API_KEY && "
"FASTVIDEO_FA4=0 pytest ./fastvideo/tests/training/Vanilla -srP"
)
@@ -210,7 +214,8 @@ def run_training_tests():
volumes={"/root/data": model_vol})
def run_training_lora_tests():
run_test(
"export HF_HOME='/root/data/.cache' && wandb login $WANDB_API_KEY && pytest ./fastvideo/tests/training/lora/test_lora_training.py -srP"
"export HF_HOME='/root/data/.cache' && wandb login $WANDB_API_KEY && "
"FASTVIDEO_FA4=0 pytest ./fastvideo/tests/training/lora/test_lora_training.py -srP"
)
@@ -220,7 +225,7 @@ def run_training_lora_tests():
secrets=[wandb_secret, ci_env_secret])
def run_training_tests_VSA():
run_test(
"wandb login $WANDB_API_KEY && pytest ./fastvideo/tests/training/VSA -srP"
"wandb login $WANDB_API_KEY && FASTVIDEO_FA4=0 pytest ./fastvideo/tests/training/VSA -srP"
)
@@ -254,7 +259,7 @@ def run_inference_lora_tests():
@app.function(gpu="L40S:2", image=image, timeout=900, secrets=[ci_env_secret])
def run_distill_dmd_tests():
run_test(
"pytest ./fastvideo/tests/training/distill/test_distill_dmd.py -vs")
"FASTVIDEO_FA4=0 pytest ./fastvideo/tests/training/distill/test_distill_dmd.py -vs")
@app.function(gpu="L40S:2",
@@ -263,7 +268,8 @@ def run_distill_dmd_tests():
secrets=[wandb_secret, ci_env_secret])
def run_self_forcing_tests():
run_test(
"wandb login $WANDB_API_KEY && pytest ./fastvideo/tests/training/self-forcing/test_self_forcing.py -vs"
"wandb login $WANDB_API_KEY && "
"FASTVIDEO_FA4=0 pytest ./fastvideo/tests/training/self-forcing/test_self_forcing.py -vs"
)
@@ -323,7 +329,8 @@ def run_dreamverse_app_tests():
volumes={"/root/data": model_vol})
def run_train_framework_tests():
run_test(
"export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/train/models ./fastvideo/tests/train/methods -vs"
"export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && "
"FASTVIDEO_FA4=0 pytest ./fastvideo/tests/train/models ./fastvideo/tests/train/methods -vs"
)
@@ -349,7 +356,8 @@ def seed_grad_norm_references():
the local command and the ``_DEVICE_MAPPINGS`` table.
"""
run_test(
"export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && FASTVIDEO_GRADNORM_UPDATE=1 pytest ./fastvideo/tests/train/methods -vs -rs"
"export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && "
"FASTVIDEO_FA4=0 FASTVIDEO_GRADNORM_UPDATE=1 pytest ./fastvideo/tests/train/methods -vs -rs"
)