Compare commits

...
26 changed files with 425 additions and 148 deletions
+97 -26
View File
@@ -14,13 +14,9 @@ on:
- ".github/workflows/pr-test.yml"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
- "csrc/**"
workflow_dispatch:
inputs:
custom_image:
description: "Custom image from this repository (default: fastvideo-dev:py3.12-latest)"
required: false
default: "fastvideo-dev:py3.12-latest"
type: string
run_encoder_test:
description: "Run encoder-test"
required: false
@@ -56,6 +52,16 @@ on:
required: false
default: false
type: boolean
run_precision_test_STA:
description: "Run precision-test-STA"
required: false
default: false
type: boolean
run_precision_test_VSA:
description: "Run precision-test-VSA"
required: false
default: false
type: boolean
run_nightly_test:
description: "Run nightly-test"
required: false
@@ -65,6 +71,7 @@ on:
env:
PYTHONUNBUFFERED: "1"
concurrency:
group: pr-test-${{ github.ref }}
cancel-in-progress: true
@@ -84,44 +91,69 @@ jobs:
training-test: ${{ steps.filter.outputs.training-test }}
training-test-VSA: ${{ steps.filter.outputs.training-test-VSA }}
inference-test-STA: ${{ steps.filter.outputs.inference-test-STA }}
precision-test-STA: ${{ steps.filter.outputs.precision-test-STA }}
precision-test-VSA: ${{ steps.filter.outputs.precision-test-VSA }}
steps:
- uses: actions/checkout@v4
- uses: dorny/paths-filter@v3
id: filter
with:
filters: |
# Define reusable path patterns
common-paths: &common-paths
- 'pyproject.toml'
- 'docker/Dockerfile.python3.12'
sta-kernel-paths: &sta-kernel-paths
- 'csrc/attn/st_attn/**'
- 'csrc/attn/setup_sta.py'
- 'csrc/attn/config_sta.py'
- 'csrc/attn/st_attn.cpp'
vsa-kernel-paths: &vsa-kernel-paths
- 'csrc/attn/vsa/**'
- 'csrc/attn/tk/**'
- 'csrc/attn/setup_vsa.py'
- 'csrc/attn/config_vsa.py'
- 'csrc/attn/vsa.cpp'
vsa-paths: &vsa-paths
- 'fastvideo/v1/**'
- *common-paths
- *vsa-kernel-paths
# Actual tests
encoder-test:
- 'fastvideo/v1/models/encoders/**'
- 'fastvideo/v1/models/loaders/**'
- 'fastvideo/v1/tests/encoders/**'
- 'pyproject.toml'
- 'docker/Dockerfile.python3.12'
- *common-paths
vae-test:
- 'fastvideo/v1/models/vaes/**'
- 'fastvideo/v1/models/loaders/**'
- 'fastvideo/v1/tests/vaes/**'
- 'pyproject.toml'
- 'docker/Dockerfile.python3.12'
- *common-paths
transformer-test:
- 'fastvideo/v1/models/dits/**'
- 'fastvideo/v1/models/loaders/**'
- 'fastvideo/v1/tests/transformers/**'
- 'fastvideo/v1/layers/**'
- 'fastvideo/v1/attention/**'
- 'pyproject.toml'
- 'docker/Dockerfile.python3.12'
- *common-paths
training-test:
- 'fastvideo/v1/**'
- 'pyproject.toml'
- 'docker/Dockerfile.python3.12'
- *common-paths
training-test-VSA:
- 'fastvideo/v1/**'
- 'pyproject.toml'
- 'docker/Dockerfile.python3.12'
- *common-paths
- *vsa-kernel-paths
inference-test-STA:
- 'fastvideo/v1/**'
- 'pyproject.toml'
- 'docker/Dockerfile.python3.12'
- *common-paths
- *sta-kernel-paths
precision-test-STA:
- *common-paths
- *sta-kernel-paths
precision-test-VSA:
- *common-paths
- *vsa-kernel-paths
encoder-test:
needs: change-filter
@@ -134,7 +166,7 @@ jobs:
gpu_type: "NVIDIA A40"
gpu_count: 1
volume_size: 100
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:py3.12-latest' }}"
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/encoders -s"
timeout_minutes: 30
secrets:
@@ -152,7 +184,7 @@ jobs:
gpu_type: "NVIDIA A40"
gpu_count: 1
volume_size: 100
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:py3.12-latest' }}"
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/vaes -s"
timeout_minutes: 30
secrets:
@@ -170,7 +202,7 @@ jobs:
gpu_type: "NVIDIA L40S"
gpu_count: 1
volume_size: 100
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:py3.12-latest' }}"
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/transformers -s"
timeout_minutes: 30
secrets:
@@ -216,7 +248,7 @@ jobs:
gpu_count: 4
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:py3.12-latest' }}"
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/v1/tests/training/Vanilla -srP"
timeout_minutes: 30
secrets:
@@ -236,7 +268,7 @@ jobs:
gpu_count: 1
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:py3.12-latest' }}"
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/v1/tests/training/VSA -srP"
timeout_minutes: 30
secrets:
@@ -256,13 +288,51 @@ jobs:
gpu_count: 1
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:py3.12-latest' }}"
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/inference/STA -srP"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
precision-test-STA:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_precision_test_STA == 'true')
uses: ./.github/workflows/runpod-test.yml
with:
job_id: "precision-test-STA"
gpu_type: "NVIDIA H100 NVL"
gpu_count: 1
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "uv pip install -e .[test] && python csrc/attn/tests/test_sta.py"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
precision-test-VSA:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_precision_test_VSA == 'true')
uses: ./.github/workflows/runpod-test.yml
with:
job_id: "precision-test-VSA"
gpu_type: "NVIDIA H100 NVL"
gpu_count: 1
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "uv pip install -e .[test] && python csrc/attn/tests/test_block_sparse.py"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
nightly-test:
if: >-
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_nightly_test == 'true')
@@ -273,7 +343,7 @@ jobs:
gpu_count: 4
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:py3.12-latest' }}"
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/v1/tests/nightly/test_e2e_overfit_single_sample.py -vs"
timeout_minutes: 30
secrets:
@@ -282,7 +352,8 @@ jobs:
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
runpod-cleanup:
needs: [encoder-test, vae-test, transformer-test, ssim-test] # Add other jobs to this list as you create them
# Add other jobs to this list as you create them
needs: [encoder-test, vae-test, transformer-test, ssim-test, training-test, training-test-VSA, inference-test-STA, precision-test-STA, precision-test-VSA]
if: ${{ always() && ((github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) || github.event_name == 'workflow_dispatch') }}
runs-on: ubuntu-latest
steps:
@@ -299,7 +370,7 @@ jobs:
- name: Cleanup all RunPod instances
env:
JOB_IDS: '["encoder-test", "vae-test", "transformer-test", "ssim-test-py3.10", "ssim-test-py3.11", "ssim-test-py3.12"]'
JOB_IDS: '["encoder-test", "vae-test", "transformer-test", "ssim-test-py3.10", "ssim-test-py3.11", "ssim-test-py3.12", "training-test", "training-test-VSA", "inference-test-STA", "precision-test-STA", "precision-test-VSA"]'
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
run: python .github/scripts/runpod_cleanup.py
+8 -2
View File
@@ -4,7 +4,7 @@
## Installation
We test our code on Pytorch 2.5.0 and CUDA>=12.4. Currently we only have implementation on H100.
We test our code on Pytorch 2.5.0 and CUDA>=12.4. Currently we only support H100/H200, because ThunderKittens uses TMA but doesn't support Blackwell yet.
First, install C++20 for ThunderKittens:
```bash
sudo apt update
@@ -53,8 +53,14 @@ out = sliding_tile_attention(q, k, v, window_size, 0, False)
## Test
```bash
python test/test_sta.py
python tests/test_sta.py # test STA
python tests/test_block_sparse.py # test VSA
```
## Benchmark
```bash
python benchmarks/bench_sta.py
```
## How Does STA Work?
We give a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
@@ -5,6 +5,7 @@ import matplotlib.pyplot as plt
import numpy as np
import torch
from st_attn import sliding_tile_attention
from triton.testing import do_bench
def flops(batch, seqlen, nheads, headdim, causal, mode="fwd"):
@@ -13,16 +14,16 @@ def flops(batch, seqlen, nheads, headdim, causal, mode="fwd"):
return f if mode == "fwd" else (2.5 * f if mode == "bwd" else 3.5 * f)
def efficiency(flop, time):
flop = flop / 1e12
time = time / 1e6
return flop / time
def compute_TFLOPS(flops, ms):
flops = flops / 1e12
ms = ms / 1e3
return flops / ms
def benchmark_attention(configurations):
results = {'fwd': defaultdict(list), 'bwd': defaultdict(list)}
for B, H, N, D, causal in configurations:
for B, H, N, D, causal, dit_seq_shape, window_size in configurations:
print("=" * 60)
print(f"Timing forward and backward pass for B={B}, H={H}, N={N}, D={D}, causal={causal}")
@@ -30,38 +31,31 @@ def benchmark_attention(configurations):
k = torch.randn(B, H, N, D, dtype=torch.bfloat16, device='cuda', requires_grad=False).contiguous()
v = torch.randn(B, H, N, D, dtype=torch.bfloat16, device='cuda', requires_grad=False).contiguous()
grad_output = torch.randn_like(q, requires_grad=False).contiguous()
# grad_output = torch.randn_like(q, requires_grad=False).contiguous()
# qg = torch.zeros_like(q, requires_grad=False, dtype=torch.float).contiguous()
# kg = torch.zeros_like(k, requires_grad=False, dtype=torch.float).contiguous()
# vg = torch.zeros_like(v, requires_grad=False, dtype=torch.float).contiguous()
qg = torch.zeros_like(q, requires_grad=False, dtype=torch.float).contiguous()
kg = torch.zeros_like(k, requires_grad=False, dtype=torch.float).contiguous()
vg = torch.zeros_like(v, requires_grad=False, dtype=torch.float).contiguous()
# Prepare for timing forward pass
start_events_fwd = [torch.cuda.Event(enable_timing=True) for _ in range(10)]
end_events_fwd = [torch.cuda.Event(enable_timing=True) for _ in range(10)]
# # Warmup for forward pass
# for _ in range(10):
# o = sliding_tile_attention(q, k, v, [[3, 6, 10]] * 24, 0, False, dit_seq_shape)
torch.cuda.empty_cache()
torch.cuda.synchronize()
# # Time the forward pass
# for i in range(10):
# start_events_fwd[i].record()
# o = sliding_tile_attention(q, k, v, [[3, 6, 10]] * 24, 0, False, dit_seq_shape)
# end_events_fwd[i].record()
ms = do_bench(lambda: sliding_tile_attention(q, k, v, [window_size] * 24, 0, False, dit_seq_shape))
# Warmup for forward pass
for _ in range(10):
o = sliding_tile_attention(q, k, v, [[3, 6, 10]] * 24, 0, False, '18x48x80')
# times_fwd = [s.elapsed_time(e) for s, e in zip(start_events_fwd, end_events_fwd)]
# time_us_fwd = np.mean(times_fwd) * 1000
# Time the forward pass
for i in range(10):
start_events_fwd[i].record()
o = sliding_tile_attention(q, k, v, [[3, 6, 10]] * 24, 0, False, '18x48x80')
end_events_fwd[i].record()
torch.cuda.synchronize()
times_fwd = [s.elapsed_time(e) for s, e in zip(start_events_fwd, end_events_fwd)]
time_us_fwd = np.mean(times_fwd) * 1000
tflops_fwd = efficiency(flops(B, N, H, D, causal, 'fwd'), time_us_fwd)
tflops_fwd = compute_TFLOPS(flops(B, N, H, D, causal, 'fwd'), ms)
results['fwd'][(D, causal)].append((N, tflops_fwd))
print(f"Average time for forward pass in us: {time_us_fwd:.2f}")
print(f"Average efficiency for forward pass in TFLOPS: {tflops_fwd}")
print(f"Average time for forward pass (ms): {ms:.2f}")
print(f"Average TFLOPS: {tflops_fwd}")
print("-" * 60)
# torch.cuda.empty_cache()
@@ -85,15 +79,14 @@ def benchmark_attention(configurations):
# times_bwd = [s.elapsed_time(e) for s, e in zip(start_events_bwd, end_events_bwd)]
# time_us_bwd = np.mean(times_bwd) * 1000
# tflops_bwd = efficiency(flops(B, N, H, D, causal, 'bwd'), time_us_bwd)
# tflops_bwd = compute_TFLOPS(flops(B, N, H, D, causal, 'bwd'), ms)
# results['bwd'][(D, causal)].append((N, tflops_bwd))
# print(f"Average time for backward pass in us: {time_us_bwd:.2f}")
# print(f"Average efficiency for backward pass in TFLOPS: {tflops_bwd}")
print("=" * 60)
# print(f"Average time for backward pass(ms): {ms:.2f}")
# print(f"Average TFLOPS: {tflops_bwd}")
# print("=" * 60)
torch.cuda.empty_cache()
torch.cuda.synchronize()
return results
@@ -124,7 +117,10 @@ def plot_results(results):
# Example list of configurations to test
configurations = [
(2, 24, 69120, 128, False),
(2, 24, 69120, 128, False, '18x48x80', [3, 6, 10]),
(2, 24, 69120, 128, True, '18x48x80', [3, 6, 10]),
(2, 24, 82944, 128, False, '36x48x48', [3, 3, 6]), # Stepvideo
(2, 24, 82944, 128, True, '36x48x48', [3, 3, 6]),
# (16, 16, 768*16, 128, False),
# (16, 16, 768*2, 128, False),
# (16, 16, 768*4, 128, False),
+31 -22
View File
@@ -4,9 +4,17 @@
#include <cooperative_groups.h>
#include <iostream>
#include <stdio.h>
#include <c10/cuda/CUDAGuard.h>
// #define CLAMP(value, min, max) ((value) < (min) ? (min) : ((value) > (max) ? (max) : (value)))
__device__ __forceinline__ int clamp_int(int value, int min, int max) {
return (value < min) ? min : ((value > max) ? max : value);
}
// #define ABS(x) ((x) < 0 ? -(x) : (x))
__device__ __forceinline__ int abs_int(int value) {
return (value < 0) ? -value : value;
}
#define CLAMP(value, min, max) ((value) < (min) ? (min) : ((value) > (max) ? (max) : (value)))
#define ABS(x) ((x) < 0 ? -(x) : (x))
constexpr int CONSUMER_WARPGROUPS = (3);
constexpr int PRODUCER_WARPGROUPS = (1);
@@ -117,16 +125,16 @@ void fwd_attend_ker(const __grid_constant__ fwd_globals<D> g) {
int qt = seq_idx / 6 / (CH * CW);
int qh = (seq_idx / 6) % (CH * CW) / CW;
int qw = (seq_idx / 6) % CW;
qt = CLAMP(qt, DT, CT-DT-1);
qh = CLAMP(qh, DH, CH-DH-1);
qw = CLAMP(qw, DW, CW-DW-1);
qt = clamp_int(qt, DT, CT-DT-1);
qh = clamp_int(qh, DH, CH-DH-1);
qw = clamp_int(qw, DW, CW-DW-1);
int count = 0;
int j = 0;
while (count < K::stages - 1) {
int kt = j / 3 / (CH * CW);
int kh = (j / 3) % (CH * CW) / CW;
int kw = (j / 3) % CW;
bool mask = (ABS(qt - kt) <= DT) && (ABS(qh - kh) <= DH) && (ABS(qw - kw) <= DW);
bool mask = (abs_int(qt - kt) <= DT) && (abs_int(qh - kh) <= DH) && (abs_int(qw - kw) <= DW);
if (mask){
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, j, 0};
tma::expect_bytes(k_smem_arrived[count], sizeof(k_tile));
@@ -167,15 +175,15 @@ void fwd_attend_ker(const __grid_constant__ fwd_globals<D> g) {
int qt = seq_idx / 6 / (CH * CW);
int qh = (seq_idx / 6) % (CH * CW) / CW;
int qw = (seq_idx / 6) % CW;
qt = CLAMP(qt, DT, CT-DT-1);
qh = CLAMP(qh, DH, CH-DH-1);
qw = CLAMP(qw, DW, CW-DW-1);
int k_t_min = CLAMP(qt-DT, 0, CT-1);
int k_t_max = CLAMP(qt+DT, 0, CT-1);
int k_h_min = CLAMP(qh-DH, 0, CH-1);
int k_h_max = CLAMP(qh+DH, 0, CH-1);
int k_w_min = CLAMP(qw-DW, 0, CW-1);
int k_w_max = CLAMP(qw+DW, 0, CW-1);
qt = clamp_int(qt, DT, CT-DT-1);
qh = clamp_int(qh, DH, CH-DH-1);
qw = clamp_int(qw, DW, CW-DW-1);
int k_t_min = clamp_int(qt-DT, 0, CT-1);
int k_t_max = clamp_int(qt+DT, 0, CT-1);
int k_h_min = clamp_int(qh-DH, 0, CH-1);
int k_h_max = clamp_int(qh+DH, 0, CH-1);
int k_w_min = clamp_int(qw-DW, 0, CW-1);
int k_w_max = clamp_int(qw+DW, 0, CW-1);
int count = 0;
for (int kt = k_t_min; kt <= k_t_max; kt++) {
for (int kh = k_h_min; kh <= k_h_max; kh++) {
@@ -234,7 +242,7 @@ void fwd_attend_ker(const __grid_constant__ fwd_globals<D> g) {
// the last three kv blocks are for text, we process them separately
kv_iters = img_kv_blocks - 1;
} else {
kv_iters = CLAMP(DT*2+1, 1, CT) * CLAMP(DH*2+1, 1, CH) * CLAMP(DW*2+1, 1, CW) * 3 - 1 ;
kv_iters = clamp_int(DT*2+1, 1, CT) * clamp_int(DH*2+1, 1, CH) * clamp_int(DW*2+1, 1, CW) * 3 - 1 ;
}
kittens::wait(qsmem_semaphore, 0);
@@ -415,8 +423,9 @@ sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o,
float* l_ptr = reinterpret_cast<float*>(l_vec.data_ptr<float>());
float* d_l = reinterpret_cast<float*>(l_ptr);
cudaDeviceSynchronize();
auto stream = at::cuda::getCurrentCUDAStream().stream();
//cudadevicesynchronize();
const c10::cuda::OptionalCUDAGuard device_guard(q.device());
const cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream();
if (head_dim == 128) {
@@ -442,8 +451,8 @@ sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o,
globals g{qg_arg, kg_arg, vg_arg, lg_arg, og_arg, static_cast<int>(seq_len), static_cast<int>(text_length), static_cast<int>(hr)};
auto mem_size = kittens::MAX_SHARED_MEMORY;
auto threads = NUM_WORKERS * kittens::WARP_THREADS;
constexpr int mem_size = kittens::MAX_SHARED_MEMORY;
int threads = NUM_WORKERS * kittens::WARP_THREADS;
if (has_text) {
// TORCH_CHECK(seq_len % (CONSUMER_WARPGROUPS*kittens::TILE_DIM*4) == 0, "sequence length must be divisible by 192");
dim3 grid_image(seq_len/(CONSUMER_WARPGROUPS*kittens::TILE_ROW_DIM<bf16>*4)-2, qo_heads, batch);
@@ -823,10 +832,10 @@ sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o,
}
CHECK_CUDA_ERROR(cudaGetLastError());
cudaStreamSynchronize(stream);
// cudaStreamSynchronize(stream);
}
return o;
cudaDeviceSynchronize();
//cudadevicesynchronize();
}
@@ -7,6 +7,7 @@ from vsa import BLOCK_M, BLOCK_N
import numpy as np
import random
import gc
def set_seed(seed: int = 42):
# Python random module
@@ -20,15 +21,6 @@ def set_seed(seed: int = 42):
torch.cuda.manual_seed(seed)
torch.cuda.manual_seed_all(seed) # if using multi-GPU
def parse_arguments():
parser = argparse.ArgumentParser(description='Benchmark Block Sparse Attention')
parser.add_argument('--batch_size', type=int, default=4, help='Batch size')
parser.add_argument('--num_heads', type=int, default=6, help='Number of heads')
parser.add_argument('--head_dim', type=int, default=128, help='Head dimension')
parser.add_argument('--topk', type=int, default=64, help='Number of kv blocks each q block attends to')
parser.add_argument('--seq_lengths', type=int, nargs='+', default=[29120], help='Sequence lengths to benchmark')
parser.add_argument('--num_iterations', type=int, default=100, help='Number of test iterations to run')
return parser.parse_args()
@torch.no_grad
def precision_metric(quant_o, fa2_o):
@@ -135,9 +127,7 @@ def generate_block_sparse_pattern(bs, h, num_q_blocks, num_kv_blocks, k, device=
return q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, block_sparse_mask
def main():
args = parse_arguments()
def main(args):
set_seed(42)
# Extract parameters
@@ -191,23 +181,36 @@ def main():
block_mask_expanded = block_sparse_mask.unsqueeze(-1).unsqueeze(-2) # [b, h, num_q_blocks, num_kv_blocks, 1, 1]
block_mask_expanded = block_mask_expanded.expand(-1, -1, -1, -1, BLOCK_M, BLOCK_N) # [b, h, num_q_blocks, num_kv_blocks, BLOCK_M, BLOCK_N]
full_mask = block_mask_expanded.permute(0, 1, 2, 4, 3, 5).reshape(batch, head, seq_len, seq_len)
q_sdpa = q.clone()
k_sdpa = k.clone()
v_sdpa = v.clone()
q.requires_grad = True
k.requires_grad = True
v.requires_grad = True
q_sdpa.requires_grad = True
k_sdpa.requires_grad = True
v_sdpa.requires_grad = True
# testing forward
o = BlockSparseAttentionFunction.apply(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num)
del q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, block_sparse_mask, block_mask_expanded
grad_o = torch.randn_like(o)
o.backward(grad_o)
# clear memory
q_sdpa = q.detach().clone()
k_sdpa = k.detach().clone()
v_sdpa = v.detach().clone()
q_sdpa.requires_grad = True
k_sdpa.requires_grad = True
v_sdpa.requires_grad = True
q.data = torch.empty(0, device=q.device)
k.data = torch.empty(0, device=k.device)
v.data = torch.empty(0, device=v.device)
torch.cuda.empty_cache()
o_sdpa = torch.nn.functional.scaled_dot_product_attention(q_sdpa, k_sdpa, v_sdpa, attn_mask=full_mask)
sim, l1, rmse = precision_metric(o, o_sdpa)
assert sim > 0.9999, f"SSIM too low: {sim}"
assert l1 < 8e-5, f"l1 too large: {l1}"
assert rmse < 2e-5, f"RMSE too large: {rmse}"
forward_metrics['sim'].append(sim)
forward_metrics['l1'].append(l1)
forward_metrics['rmse'].append(rmse)
@@ -215,52 +218,72 @@ def main():
print(f"block_sparse_attention_fwd vs torch.nn.functional.scaled_dot_product_attention:\nsim: {sim}, l1: {l1}, rmse: {rmse}")
# test backward
grad_o = torch.randn_like(o)
o.backward(grad_o)
o_sdpa.backward(grad_o)
sim, l1, rmse = precision_metric(q.grad, q_sdpa.grad)
# Error bounds collected on H100
assert sim > 0.9999, f"SSIM too low: {sim}"
assert l1 < 4e-3, f"l1 too large: {l1}"
assert rmse < 3e-4, f"RMSE too large: {rmse}"
grad_q_metrics['sim'].append(sim)
grad_q_metrics['l1'].append(l1)
grad_q_metrics['rmse'].append(rmse)
print(f"block_sparse_attention_bwd vs torch.nn.functional.scaled_dot_product_attention grad_q:\nsim: {sim}, l1: {l1}, rmse: {rmse}")
sim, l1, rmse = precision_metric(k.grad, k_sdpa.grad)
assert sim > 0.9999, f"SSIM too low: {sim}"
assert l1 < 4e-3, f"l1 too large: {l1}"
assert rmse < 2e-4, f"RMSE too large: {rmse}"
grad_k_metrics['sim'].append(sim)
grad_k_metrics['l1'].append(l1)
grad_k_metrics['rmse'].append(rmse)
print(f"block_sparse_attention_bwd vs torch.nn.functional.scaled_dot_product_attention grad_k:\nsim: {sim}, l1: {l1}, rmse: {rmse}")
sim, l1, rmse = precision_metric(v.grad, v_sdpa.grad)
assert sim > 0.9999, f"SSIM too low: {sim}"
assert l1 < 1e-4, f"l1 too large: {l1}"
assert rmse < 2e-5, f"RMSE too large: {rmse}"
grad_v_metrics['sim'].append(sim)
grad_v_metrics['l1'].append(l1)
grad_v_metrics['rmse'].append(rmse)
print(f"block_sparse_attention_bwd vs torch.nn.functional.scaled_dot_product_attention grad_v:\nsim: {sim}, l1: {l1}, rmse: {rmse}")
del o, o_sdpa, grad_o, q_sdpa, k_sdpa, v_sdpa
gc.collect()
torch.cuda.empty_cache()
# Print summary statistics if multiple iterations were run
if num_iterations > 1:
print("\n" + "="*50)
print(f"Summary Statistics (over {num_iterations} iterations):")
print("\nForward metrics:")
print(f"Similarity: mean={np.mean(forward_metrics['sim']):.6f}, std={np.std(forward_metrics['sim']):.6f}")
print(f"L1 error: mean={np.mean(forward_metrics['l1']):.6f}, std={np.std(forward_metrics['l1']):.6f}")
print(f"RMSE: mean={np.mean(forward_metrics['rmse']):.6f}, std={np.std(forward_metrics['rmse']):.6f}")
print(f"Similarity: mean={np.mean(forward_metrics['sim']):.6f}, std={np.std(forward_metrics['sim']):.6f}, min={np.min(forward_metrics['sim']):.6f}")
print(f"L1 error: mean={np.mean(forward_metrics['l1']):.6f}, std={np.std(forward_metrics['l1']):.6f}, max={np.max(forward_metrics['l1']):.6f}")
print(f"RMSE: mean={np.mean(forward_metrics['rmse']):.6f}, std={np.std(forward_metrics['rmse']):.6f}, max={np.max(forward_metrics['rmse']):.6f}")
print("\nGradient Q metrics:")
print(f"Similarity: mean={np.mean(grad_q_metrics['sim']):.6f}, std={np.std(grad_q_metrics['sim']):.6f}")
print(f"L1 error: mean={np.mean(grad_q_metrics['l1']):.6f}, std={np.std(grad_q_metrics['l1']):.6f}")
print(f"RMSE: mean={np.mean(grad_q_metrics['rmse']):.6f}, std={np.std(grad_q_metrics['rmse']):.6f}")
print(f"Similarity: mean={np.mean(grad_q_metrics['sim']):.6f}, std={np.std(grad_q_metrics['sim']):.6f}, min={np.min(grad_q_metrics['sim']):.6f}")
print(f"L1 error: mean={np.mean(grad_q_metrics['l1']):.6f}, std={np.std(grad_q_metrics['l1']):.6f}, max={np.max(grad_q_metrics['l1']):.6f}")
print(f"RMSE: mean={np.mean(grad_q_metrics['rmse']):.6f}, std={np.std(grad_q_metrics['rmse']):.6f}, max={np.max(grad_q_metrics['rmse']):.6f}")
print("\nGradient K metrics:")
print(f"Similarity: mean={np.mean(grad_k_metrics['sim']):.6f}, std={np.std(grad_k_metrics['sim']):.6f}")
print(f"L1 error: mean={np.mean(grad_k_metrics['l1']):.6f}, std={np.std(grad_k_metrics['l1']):.6f}")
print(f"RMSE: mean={np.mean(grad_k_metrics['rmse']):.6f}, std={np.std(grad_k_metrics['rmse']):.6f}")
print(f"Similarity: mean={np.mean(grad_k_metrics['sim']):.6f}, std={np.std(grad_k_metrics['sim']):.6f}, min={np.min(grad_k_metrics['sim']):.6f}")
print(f"L1 error: mean={np.mean(grad_k_metrics['l1']):.6f}, std={np.std(grad_k_metrics['l1']):.6f}, max={np.max(grad_k_metrics['l1']):.6f}")
print(f"RMSE: mean={np.mean(grad_k_metrics['rmse']):.6f}, std={np.std(grad_k_metrics['rmse']):.6f}, max={np.max(grad_k_metrics['rmse']):.6f}")
print("\nGradient V metrics:")
print(f"Similarity: mean={np.mean(grad_v_metrics['sim']):.6f}, std={np.std(grad_v_metrics['sim']):.6f}")
print(f"L1 error: mean={np.mean(grad_v_metrics['l1']):.6f}, std={np.std(grad_v_metrics['l1']):.6f}")
print(f"RMSE: mean={np.mean(grad_v_metrics['rmse']):.6f}, std={np.std(grad_v_metrics['rmse']):.6f}")
print(f"Similarity: mean={np.mean(grad_v_metrics['sim']):.6f}, std={np.std(grad_v_metrics['sim']):.6f}, min={np.min(grad_v_metrics['sim']):.6f}")
print(f"L1 error: mean={np.mean(grad_v_metrics['l1']):.6f}, std={np.std(grad_v_metrics['l1']):.6f}, max={np.max(grad_v_metrics['l1']):.6f}")
print(f"RMSE: mean={np.mean(grad_v_metrics['rmse']):.6f}, std={np.std(grad_v_metrics['rmse']):.6f}, max={np.max(grad_v_metrics['rmse']):.6f}")
if __name__ == "__main__":
main()
parser = argparse.ArgumentParser(description='Benchmark Block Sparse Attention')
parser.add_argument('--batch_size', type=int, default=4, help='Batch size')
parser.add_argument('--num_heads', type=int, default=6, help='Number of heads')
parser.add_argument('--head_dim', type=int, default=128, help='Head dimension')
parser.add_argument('--topk', type=int, default=64, help='Number of kv blocks each q block attends to')
parser.add_argument('--seq_lengths', type=int, nargs='+', default=[29120], help='Sequence lengths to benchmark')
parser.add_argument('--num_iterations', type=int, default=50, help='Number of test iterations to run')
args = parser.parse_args()
main(args)
@@ -81,5 +81,7 @@ std = 10
# Run correctness check directly
results = check_correctness(b, h, n, d, causal, mean, std, error_mode='output')
assert results['TK vs FLEX']['avg_diff'] < 3e-6, f"Average difference: {results['TK vs FLEX']['avg_diff']} is too large"
assert results['TK vs FLEX']['max_diff'] < 4e-2, f"Maximum difference: {results['TK vs FLEX']['max_diff']} is too large"
print(f"Average difference: {results['TK vs FLEX']['avg_diff']}")
print(f"Maximum difference: {results['TK vs FLEX']['max_diff']}")
+23 -19
View File
@@ -3,6 +3,8 @@
#include "kittens.cuh"
#include <cooperative_groups.h>
#include <iostream>
#include <c10/cuda/CUDAGuard.h>
using namespace kittens;
namespace cg = cooperative_groups;
@@ -940,8 +942,9 @@ block_sparse_attention_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v
float* l_ptr = reinterpret_cast<float*>(l_vec.data_ptr<float>());
float* d_l = reinterpret_cast<float*>(l_ptr);
cudaDeviceSynchronize();
auto stream = at::cuda::getCurrentCUDAStream().stream();
//cudadevicesynchronize();
const c10::cuda::OptionalCUDAGuard device_guard(q.device());
const cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream();
if (head_dim == 64) {
using q_tile = st_bf<fwd_attend_ker_tile_dims<64>::qo_height, fwd_attend_ker_tile_dims<64>::tile_width>;
@@ -966,7 +969,7 @@ block_sparse_attention_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v
globals g{qg_arg, kg_arg, vg_arg, lg_arg, og_arg, static_cast<int>(seq_len), static_cast<int>(hr), static_cast<int>(max_kv_blocks_per_q), reinterpret_cast<int32_t*>(q2k_block_sparse_index.data_ptr()), reinterpret_cast<int32_t*>(q2k_block_sparse_num.data_ptr())};
auto mem_size = 54000;
constexpr int mem_size = 54000;
dim3 grid(seq_len/(64), qo_heads, batch);
@@ -979,7 +982,7 @@ block_sparse_attention_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v
fwd_attend_ker<64><<<grid, (128), mem_size, stream>>>(g);
CHECK_CUDA_ERROR(cudaGetLastError());
cudaStreamSynchronize(stream);
// cudaStreamSynchronize(stream);
}
if (head_dim == 128) {
@@ -1005,7 +1008,7 @@ block_sparse_attention_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v
globals g{qg_arg, kg_arg, vg_arg, lg_arg, og_arg, static_cast<int>(seq_len), static_cast<int>(hr), static_cast<int>(max_kv_blocks_per_q), reinterpret_cast<int32_t*>(q2k_block_sparse_index.data_ptr()), reinterpret_cast<int32_t*>(q2k_block_sparse_num.data_ptr())};
auto mem_size = 54000;
constexpr int mem_size = 54000;
dim3 grid(seq_len/(64), qo_heads, batch);
@@ -1018,11 +1021,11 @@ block_sparse_attention_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v
fwd_attend_ker<128><<<grid, (128), mem_size, stream>>>(g);
CHECK_CUDA_ERROR(cudaGetLastError());
cudaStreamSynchronize(stream);
// cudaStreamSynchronize(stream);
}
return {o, l_vec};
cudaDeviceSynchronize();
//cudadevicesynchronize();
}
std::vector<torch::Tensor>
@@ -1132,13 +1135,14 @@ block_sparse_attention_backward(torch::Tensor q,
float* d_kg = reinterpret_cast<float*>(kg_ptr);
float* d_vg = reinterpret_cast<float*>(vg_ptr);
auto mem_size = kittens::MAX_SHARED_MEMORY;
auto threads = 4 * kittens::WARP_THREADS;
constexpr int mem_size = kittens::MAX_SHARED_MEMORY;
int threads = 4 * kittens::WARP_THREADS;
cudaDeviceSynchronize();
auto stream = at::cuda::getCurrentCUDAStream().stream();
//cudadevicesynchronize();
const c10::cuda::OptionalCUDAGuard device_guard(q.device());
const cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream();
cudaStreamSynchronize(stream);
// cudaStreamSynchronize(stream);
// TORCH_CHECK(seq_len % (4*kittens::TILE_DIM*4) == 0, "sequence length must be divisible by 256");
dim3 grid_bwd(seq_len/(4*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
@@ -1222,7 +1226,7 @@ block_sparse_attention_backward(torch::Tensor q,
dim3 grid_bwd_2(seq_len/64, qo_heads, batch);
threads = 128;
cudaDeviceSynchronize();
//cudadevicesynchronize();
{
cudaFuncSetAttribute(
@@ -1240,8 +1244,8 @@ block_sparse_attention_backward(torch::Tensor q,
}
// CHECK_CUDA_ERROR(cudaGetLastError());
cudaStreamSynchronize(stream);
cudaDeviceSynchronize();
// cudaStreamSynchronize(stream);
//cudadevicesynchronize();
// const auto kernel_end = std::chrono::high_resolution_clock::now();
// std::cout << "Kernel Time: " << std::chrono::duration_cast<std::chrono::microseconds>(kernel_end - start).count() << "us" << std::endl;
// std::cout << "---" << std::endl;
@@ -1326,7 +1330,7 @@ block_sparse_attention_backward(torch::Tensor q,
dim3 grid_bwd_2(seq_len/64, qo_heads, batch);
threads = 128;
cudaDeviceSynchronize();
//cudadevicesynchronize();
{
cudaFuncSetAttribute(
@@ -1338,10 +1342,10 @@ block_sparse_attention_backward(torch::Tensor q,
bwd_attend_ker<128><<<grid_bwd_2, threads, 113000, stream>>>(bwd_global);
}
cudaStreamSynchronize(stream);
cudaDeviceSynchronize();
// cudaStreamSynchronize(stream);
//cudadevicesynchronize();
}
return {qg, kg, vg};
cudaDeviceSynchronize();
//cudadevicesynchronize();
}
+1
View File
@@ -12,6 +12,7 @@ class DiTArchConfig(ArchConfig):
_fsdp_shard_conditions: list = field(default_factory=list)
_compile_conditions: list = field(default_factory=list)
_param_names_mapping: dict = field(default_factory=dict)
_reverse_param_names_mapping: dict = field(default_factory=dict)
_lora_param_names_mapping: dict = field(default_factory=dict)
_supported_attention_backends: Tuple[AttentionBackendEnum, ...] = (
AttentionBackendEnum.SLIDING_TILE_ATTN, AttentionBackendEnum.SAGE_ATTN,
@@ -147,6 +147,9 @@ class HunyuanVideoArchConfig(DiTArchConfig):
r"final_layer.linear.\1",
})
# Reverse mapping for saving checkpoints: training -> diffusers
_reverse_param_names_mapping: dict = field(default_factory=lambda: {})
patch_size: int = 2
patch_size_t: int = 1
in_channels: int = 16
+5 -1
View File
@@ -49,9 +49,13 @@ class WanVideoArchConfig(DiTArchConfig):
r"blocks.\1.ffn.fc_in.\2",
r"^blocks\.(\d+)\.ffn\.net\.2\.(.*)$":
r"blocks.\1.ffn.fc_out.\2",
r"blocks\.(\d+)\.norm2\.(.*)$":
r"^blocks\.(\d+)\.norm2\.(.*)$":
r"blocks.\1.self_attn_residual_norm.norm.\2",
})
# Reverse mapping for saving checkpoints: training -> diffusers
_reverse_param_names_mapping: dict = field(default_factory=lambda: {})
# Some LoRA adapters use the original official layer names instead of hf layer names,
# so apply this before the param_names_mapping
_lora_param_names_mapping: dict = field(
+2
View File
@@ -14,6 +14,7 @@ class BaseDiT(nn.Module, ABC):
_fsdp_shard_conditions: list = []
_compile_conditions: list = []
_param_names_mapping: dict
_reverse_param_names_mapping: dict
hidden_size: int
num_attention_heads: int
num_channels_latents: int
@@ -78,6 +79,7 @@ class CachableDiT(BaseDiT):
# These are required class attributes that should be overridden by concrete implementations
_fsdp_shard_conditions = []
_param_names_mapping = {}
_reverse_param_names_mapping = {}
_lora_param_names_mapping: dict = {}
# Ensure these instance attributes are properly defined in subclasses
hidden_size: int
+2
View File
@@ -442,6 +442,8 @@ class HunyuanVideoTransformer3DModel(CachableDiT):
_supported_attention_backends = HunyuanVideoConfig(
)._supported_attention_backends
_param_names_mapping = HunyuanVideoConfig()._param_names_mapping
_reverse_param_names_mapping = HunyuanVideoConfig(
)._reverse_param_names_mapping
_lora_param_names_mapping = HunyuanVideoConfig()._lora_param_names_mapping
def __init__(self, config: HunyuanVideoConfig, hf_config: dict[str, Any]):
+2
View File
@@ -460,6 +460,8 @@ class StepVideoModel(BaseDiT):
# lambda n, m: "pos_embed" in n # If needed for the patch embedding.
]
_param_names_mapping = StepVideoConfig()._param_names_mapping
_reverse_param_names_mapping = StepVideoConfig(
)._reverse_param_names_mapping
_lora_param_names_mapping = StepVideoConfig()._lora_param_names_mapping
_supported_attention_backends = StepVideoConfig(
)._supported_attention_backends
+1
View File
@@ -518,6 +518,7 @@ class WanTransformer3DModel(CachableDiT):
_supported_attention_backends = WanVideoConfig(
)._supported_attention_backends
_param_names_mapping = WanVideoConfig()._param_names_mapping
_reverse_param_names_mapping = WanVideoConfig()._reverse_param_names_mapping
_lora_param_names_mapping = WanVideoConfig()._lora_param_names_mapping
def __init__(self, config: WanVideoConfig, hf_config: dict[str,
+6 -1
View File
@@ -222,10 +222,14 @@ def load_model_from_full_model_state_dict(
used_keys = set()
sharded_sd = {}
to_merge_params: DefaultDict[str, Dict[Any, Any]] = defaultdict(dict)
reverse_param_names_mapping = {}
assert param_names_mapping is not None
for source_param_name, full_tensor in full_sd_iterator:
assert param_names_mapping is not None
target_param_name, merge_index, num_params_to_merge = param_names_mapping(
source_param_name)
reverse_param_names_mapping[target_param_name] = (source_param_name,
merge_index,
num_params_to_merge)
used_keys.add(target_param_name)
if merge_index is not None:
to_merge_params[target_param_name][merge_index] = full_tensor
@@ -260,6 +264,7 @@ def load_model_from_full_model_state_dict(
sharded_tensor = sharded_tensor.cpu()
sharded_sd[target_param_name] = nn.Parameter(sharded_tensor)
model._reverse_param_names_mapping = reverse_param_names_mapping
unused_keys = set(meta_sd.keys()) - used_keys
if unused_keys:
logger.warning("Found new parameters in meta state dict: %s",
@@ -0,0 +1,51 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
DATA_DIR="data/crush-smol_processed_main_t2v/latents/combined_parquet_dataset"
VALIDATION_DIR="data/crush-smol_processed_main_t2v/latents/validation_parquet_dataset"
NUM_GPUS=4
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
fastvideo/v1/training/wan_training_pipeline.py\
--model_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--inference_mode False\
--pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--data_path "$DATA_DIR"\
--validation_preprocessed_path "$VALIDATION_DIR"\
--train_batch_size=1 \
--num_latent_t 8 \
--sp_size 4 \
--tp_size 4 \
--hsdp_replicate_dim 1 \
--hsdp_shard_dim 4 \
--num_gpus $NUM_GPUS \
--train_sp_batch_size 1\
--dataloader_num_workers 1\
--gradient_accumulation_steps=8 \
--max_train_steps=5000 \
--learning_rate=1e-5\
--mixed_precision="bf16"\
--checkpointing_steps=6000 \
--validation_steps 50\
--validation_sampling_steps "50" \
--log_validation \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--output_dir="$DATA_DIR/outputs/wan_finetune"\
--tracker_project_name wan_finetune \
--num_height 480 \
--num_width 832 \
--num_frames 77 \
--validation_guidance_scale "1.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--weight_decay 0.01 \
--not_apply_cfg_solver \
--dit_precision "fp32" \
--max_grad_norm 1.0
@@ -0,0 +1,24 @@
# export WANDB_MODE="offline"
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_TYPE="wan"
DATA_MERGE_PATH="data/crush-smol/merge.txt"
OUTPUT_DIR="data/crush-smol_processed_main_t2v/latents"
VALIDATION_PATH="examples/training/finetune/wan_t2v_1.3b/crush_smol/validation.json"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/v1/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 8 \
--max_height 480 \
--max_width 832 \
--num_frames 77 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--model_type $MODEL_TYPE \
--train_fps 16 \
--validation_dataset_file $VALIDATION_PATH \
--samples_per_file 8 \
--flush_frequency 8 \
--preprocess_task "t2v"
@@ -116,7 +116,7 @@ def test_distributed_training():
'avg_step_time': 1.0,
'grad_norm': 0.2,
'step_time': 0.5,
'train_loss': 0.001
'train_loss': 0.0025
}
failures = []
+75 -4
View File
@@ -8,7 +8,6 @@ from typing import Any, Dict, List, Optional, Tuple, Union
import torch
import torch.distributed as dist
import torch.distributed.checkpoint as dcp
import torch.distributed.checkpoint.stateful
from einops import rearrange
from safetensors.torch import save_file
@@ -154,13 +153,20 @@ def save_checkpoint(transformer,
if rank == 0:
# Save model weights (consolidated)
weight_path = os.path.join(save_dir,
transformer_save_dir = os.path.join(save_dir, "transformer")
os.makedirs(transformer_save_dir, exist_ok=True)
weight_path = os.path.join(transformer_save_dir,
"diffusion_pytorch_model.safetensors")
logger.info("rank: %s, saving consolidated checkpoint to %s",
rank,
weight_path,
local_main_process_only=False)
save_file(cpu_state, weight_path)
# 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)
logger.info("rank: %s, consolidated checkpoint saved to %s",
rank,
weight_path,
@@ -170,7 +176,7 @@ def save_checkpoint(transformer,
config_dict = transformer.hf_config
if "dtype" in config_dict:
del config_dict["dtype"] # TODO
config_path = os.path.join(save_dir, "config.json")
config_path = os.path.join(transformer_save_dir, "config.json")
# save dict as json
with open(config_path, "w") as f:
json.dump(config_dict, f, indent=4)
@@ -479,3 +485,68 @@ def _has_foreach_support(tensors: List[torch.Tensor],
device: torch.device) -> bool:
return _device_has_foreach_support(device) and all(
t is None or type(t) in [torch.Tensor] for t in tensors)
def convert_training_to_diffusers_format(state_dict: Dict[str, Any],
transformer) -> Dict[str, Any]:
"""
Convert training format state dict to diffusers format using reverse_param_names_mapping.
Args:
state_dict: State dict in training format
transformer: Transformer model object with _reverse_param_names_mapping
Returns:
State dict in diffusers format
"""
new_state_dict = {}
# Get the reverse mapping from the transformer
reverse_param_names_mapping = transformer._reverse_param_names_mapping
assert reverse_param_names_mapping != {}, "reverse_param_names_mapping is empty"
# Group parameters that need to be split (merged parameters)
merge_groups: Dict[str, List[Tuple[str, int, int]]] = {}
# First pass: collect all merge groups
for training_key, (
diffusers_key, merge_index,
num_params_to_merge) in reverse_param_names_mapping.items():
if merge_index is not None:
# This is a merged parameter that needs to be split
if training_key not in merge_groups:
merge_groups[training_key] = []
merge_groups[training_key].append(
(diffusers_key, merge_index, num_params_to_merge))
# Second pass: handle merged parameters by splitting them
used_keys = set()
for training_key, splits in merge_groups.items():
if training_key in state_dict:
v = state_dict[training_key]
# Sort by merge_index to ensure correct order
splits.sort(key=lambda x: x[1])
total = splits[0][2]
split_size = v.shape[0] // total
split_tensors = torch.split(v, split_size, dim=0)
for diffusers_key, split_index, _ in splits:
new_state_dict[diffusers_key] = split_tensors[split_index]
used_keys.add(training_key)
# Third pass: handle regular parameters (direct mappings)
for training_key, v in state_dict.items():
if training_key in used_keys:
continue
if training_key in reverse_param_names_mapping:
diffusers_key, merge_index, _ = reverse_param_names_mapping[
training_key]
if merge_index is None:
# Direct mapping
new_state_dict[diffusers_key] = v
else:
# No mapping found, keep as is
new_state_dict[training_key] = v
return new_state_dict
+1 -1
View File
@@ -19,7 +19,7 @@ dependencies = [
# Machine Learning & Transformers
"transformers>=4.46.1", "tokenizers>=0.20.1", "sentencepiece==0.2.0",
"timm==1.0.11", "peft==0.13.2", "diffusers>=0.33.1", "bitsandbytes",
"timm==1.0.11", "peft==0.15.0", "diffusers>=0.33.1", "bitsandbytes",
"torch==2.7.1", "torchvision",
# Acceleration & Optimization