diff --git a/.github/workflows/window-process.yml b/.github/workflows/window-process.yml new file mode 100644 index 000000000..874811e2f --- /dev/null +++ b/.github/workflows/window-process.yml @@ -0,0 +1,68 @@ +# Correctness checks for the fused window kernels that need no GPU. +# +# The kernels perform no arithmetic on tensor values: each is a gather whose +# whole logic is the computation of input_offset from blockIdx. reference.py +# transcribes that to PyTorch, so the index math can be checked on a CPU runner. +# torch.utils.hipify is pure Python for the same reason, so the ROCm/HIP +# translation can be verified here too. +# +# The parity tests against the compiled extension (unit_test.py, +# test_model_parity.py) need a CUDA device, so they assert nothing on this +# runner. They are still executed, to catch the case where they fail to import +# or error out instead of skipping. + +name: window process + +on: + push: + paths: + - 'kernels/window_process/**' + - '.github/workflows/window-process.yml' + pull_request: + paths: + - 'kernels/window_process/**' + - '.github/workflows/window-process.yml' + workflow_dispatch: + +jobs: + cpu-checks: + runs-on: ubuntu-latest + strategy: + fail-fast: false + matrix: + include: + - python-version: '3.9' + torch-version: '2.8.0' + - python-version: '3.11' + torch-version: '2.13.0' + steps: + - uses: actions/checkout@v4 + + - uses: actions/setup-python@v5 + with: + python-version: ${{ matrix.python-version }} + + - name: Install PyTorch (CPU build) + # Pinned, because the two legs are not interchangeable: hipify rewrites + # the stream getter on 2.8 and leaves it on 2.13, and hipify_check.py is + # written to accept either. Resolving whatever is newest would silently + # stop covering one of the two behaviours. + run: | + python -m pip install --upgrade pip + python -m pip install numpy + python -m pip install torch==${{ matrix.torch-version }} \ + --index-url https://download.pytorch.org/whl/cpu + + - name: Index math + working-directory: kernels/window_process + run: python -m unittest test_index_math -v + + - name: CUDA to HIP translation + working-directory: kernels/window_process + run: python hipify_check.py + + - name: GPU test files import and skip cleanly + working-directory: kernels/window_process + run: | + python unit_test.py + python test_model_parity.py diff --git a/kernels/window_process/README.md b/kernels/window_process/README.md new file mode 100644 index 000000000..f3d16eda9 --- /dev/null +++ b/kernels/window_process/README.md @@ -0,0 +1,262 @@ +# Fused window process kernels + +`torch.roll` + `window_partition`, and its inverse, fused into a single pass. + +Both are pure data movement, so the eager path materialises the tensor twice — +`torch.roll` writes a copy, then `window_partition` permutes and calls +`.contiguous()` for another — where the fused kernel writes once. Peak memory +drops with the copy that is no longer made. + +Enable with `--fused_window_process` (see `get_started.md`), or by passing +`fused_window_process=True` to the model. + +## Build + +```bash +cd kernels/window_process +python setup.py install +``` + +The same command builds on ROCm; see [Building on AMD](#building-on-amd). + +## Usage + +```python +from window_process import WindowProcess, WindowProcessReverse + +# (B, H, W, C) -> (B * nH * nW, window_h, window_w, C) +windows = WindowProcess.apply(x, B, H, W, C, -shift_size, window_size) + +# and back +x = WindowProcessReverse.apply(windows, B, H, W, C, shift_size, window_size) +``` + +`shift_size` and `window_size` accept an `int` (isotropic) or an `(h, w)` pair: + +```python +windows = WindowProcess.apply(x, B, H, W, C, (-2, -8), (4, 16)) +``` + +`shift_size` is the negated `torch.roll` shift on the partition path, which is +why the call sites pass `-shift_size` forward and `+shift_size` back. + +## Constraints + +| | | +|---|---| +| dtype | float64, float32, float16, bfloat16 | +| layout | contiguous, channels last in memory: `(..., C)` | +| tiling | `H % window_h == 0` and `W % window_w == 0` | +| shape | `tensor.numel() == B * H * W * C` | +| shift | `abs(shift) < window size`, per axis | +| size | `B * H * W * C < 2^31` — offsets are computed in `int` | +| grid | `B * nH * nW <= 65535` and `H <= 65535` — CUDA `grid.z` and `grid.y` | + +Shape, tiling, shift and size are checked with `TORCH_CHECK` before the launch; +device and contiguity by `CHECK_INPUT`. The channel-last layout is the shape +contract itself and is not separately verified. + +Four of these are new guards over what was previously undetected: violating the +tiling, shape or size rules used to read the wrong element, or read out of +bounds, with no error raised anywhere, and exceeding the grid bounds used to +fail the launch asynchronously somewhere else. The `shift` row is different — +the kernels compute a larger roll correctly, for `window <= |shift| <= H` per +axis, and past that the operand of the modulo stops being non-negative and wraps +into a silently wrong read rather than an out-of-bounds one. It is enforced +because it is the contract the model already documents, not because it used to +break. dtype and layout were +already reported by the dispatch and by `CHECK_CONTIGUOUS`. + +Note that `nH != nW` follows from `H != W` alone: with a square window, +`nH = H / window_size` and `nW = W / window_size`. A non-square window is not +required to reach that case, and `img_size` is documented as `int | tuple(int)`. + +## Tests + +| file | needs a GPU | what it covers | +|---|---|---| +| `test_index_math.py` | no | the index arithmetic of all four kernels, transcribed to PyTorch in `reference.py` | +| `hipify_check.py` | no | CUDA → HIP translation is complete, and every launch carries a stream | +| `unit_test.py` | yes | parity with the eager path across dtype × shape × shift, forward and backward | +| `test_model_parity.py` | yes | a `SwinTransformer` gives identical logits and gradients either way | +| `benchmark.py` | yes | wall time and peak memory, fused vs eager | + +```bash +python test_index_math.py # runs anywhere +python hipify_check.py # runs anywhere +python unit_test.py # skips without a GPU or the built extension +python test_model_parity.py +python benchmark.py --iters 200 +``` + +The first two run in CI on a CPU runner, on every push that touches this +directory. `unit_test.py` and `test_model_parity.py` are executed there too, but +only to confirm they skip cleanly rather than error; `benchmark.py` is not, since +it exits non-zero without a device. The two CI legs pin different torch +versions, one that rewrites the stream getter and one that does not, so both +hipify behaviours are exercised on every run. + +The kernels perform no arithmetic on the values, only gathers, so parity is +asserted with `torch.equal` on every dtype including float16 and bfloat16: any +deviation is an indexing error rather than a rounding one. For the same reason +`gradcheck` adds nothing, even though float64 does dispatch. + +`reference.py` transcribes the kernels' `input_offset` computation to vectorised +PyTorch. That is not an approximation of the kernels — the offset computation is +their entire logic — which is what makes them testable on a CPU, with no +compiled extension and no GPU of either vendor. The two agree wherever the +operand of a modulo stays non-negative, which every input the launcher accepts +guarantees; outside that domain Python floors where the kernel wraps, and the +transcription is the more forgiving of the two. + +## Validation + +Parity is asserted against the eager path on an **RTX 3080** (torch +`2.8.0+cu129`, CUDA 12.9) and an **AMD Instinct MI300X** (gfx942, torch +`2.12.0+rocm7.14.0`): `unit_test.py` 15/15 test methods and +`test_model_parity.py` 4/4 on both, bit-exact, over square and non-square +images and windows including coprime `nH`/`nW`. `compute-sanitizer` (memcheck, +initcheck, synccheck) is clean on the RTX 3080 across all four kernels, forward +and backward; ROCm has no equivalent run. The MI300X run covers all 15 entries +in `unit_test.py`'s shape list; two of them were added after the RTX 3080 run +and have not been executed on NVIDIA. + +## Performance + +`benchmark.py`, forward + backward, batch 192, 200 iterations after 10 warm-up. +The `partition` direction is shown. The `merge` direction's fused time tracks it +within 6% at the configurations listed here on the RTX 3080, and within 1% at +every point of the full matrix on the MI300X. `benchmark.py` prints that matrix, +both directions and all four configurations. + +**RTX 3080 (10 GB)** + +| config | dtype | eager ms | fused ms | speedup | eager MiB | fused MiB | +|---|---|---:|---:|---:|---:|---:| +| stage 1 56×56 w7 | float32 | 7.137 | 2.398 | 2.98× | 1543.5 | 882.0 | +| stage 1 56×56 w7 | float16 | 5.223 | 2.083 | 2.51× | 771.8 | 441.0 | +| stage 1 56×56 w7 | bfloat16 | 5.074 | 1.948 | 2.60× | 771.8 | 441.0 | +| stage 2 28×28 w7 | float32 | 3.711 | 1.156 | 3.21× | 771.8 | 441.0 | +| stage 2 28×28 w7 | float16 | 2.552 | 0.700 | 3.65× | 392.0 | 224.0 | +| stage 2 28×28 w7 | bfloat16 | 2.783 | 0.726 | 3.84× | 392.0 | 224.0 | + +Non-square configurations (32×16 with a square window, 16×64 with a 4×16 window) +land in the same 2.5–3.0× band. + +**AMD Instinct MI300X** + +| config | dtype | eager ms | fused ms | speedup | eager MiB | fused MiB | +|---|---|---:|---:|---:|---:|---:| +| stage 1 56×56 w7 | float32 | 0.911 | 0.598 | 1.52× | 1543.5 | 882.0 | +| stage 1 56×56 w7 | float16 | 0.652 | 0.588 | 1.11× | 771.8 | 441.0 | +| stage 1 56×56 w7 | bfloat16 | 0.646 | 0.589 | 1.10× | 771.8 | 441.0 | +| stage 2 28×28 w7 | float32 | 0.449 | 0.188 | 2.38× | 771.8 | 441.0 | +| stage 2 28×28 w7 | float16 | 0.323 | 0.152 | 2.12× | 392.0 | 224.0 | +| stage 2 28×28 w7 | bfloat16 | 0.320 | 0.152 | 2.10× | 392.0 | 224.0 | + +Non-square configurations land in a 1.03–1.44× band. + +Memory is identical across vendors, as it must be — the allocations are the +same. Time is not. The eager baseline is what moved: 7.137 ms at stage 1 on the +3080 against 0.911 ms here, so both paths compress toward a floor and the ratio +closes with them. The fused rows also stop scaling with dtype on the MI300X +(0.598 ms float32 against 0.588 float16, for half the bytes) while the eager +path still does, so at these sizes the fused kernel is no longer bandwidth bound +on CDNA. Block width is not the cause — that is measured below. + +`bfloat16` is dispatched directly. The upstream kernel raises on a `bfloat16` +input, so the fused path previously needed an fp32 round trip. + +## Building on AMD + +`setup.py` needs no change. `CUDAExtension` hipifies its own sources when +`torch.version.hip` is set, substitutes `hipcc`, and derives `--offload-arch` +from `PYTORCH_ROCM_ARCH`: + +```bash +PYTORCH_ROCM_ARCH=gfx942 python setup.py install +``` + +On AMD's ROCm 7.14 images that command fails until the SDK headers are in +place; the second bullet below has the two `export`s and the one package it +needs, and they have to come first. + +Do not pass `--offload-arch` through `extra_compile_args`: PyTorch's +`_get_rocm_arch_flags()` skips its own detection as soon as it sees one, which +pins the build to a single GPU. + +Two things are worth knowing before building: + +- **`__ldg` has no HIP overload for `c10::Half`**, and half is in the dispatch. + Reads go through the `SWIN_WP_LDG` macro for that reason: on NVIDIA it expands + to the same `__ldg(ptr)` tokens as before, on AMD to a plain load. The hint is + advisory on both, so no result changes. `hipify_check.py` cannot catch this + and says so — it is a symbol-level check, and a symbol can translate cleanly + and still not compile. + +- **AMD's ROCm 7.14 PyTorch images ship the SDK as pip wheels** with no + `/opt/rocm` tree, and those wheels carry no rocThrust, hipSPARSE, hipBLAS, + hipBLASLt or hipSOLVER headers — which torch's own headers include when + compiling device code. In that state *no* PyTorch HIP extension compiles; a + file whose only content is `#include ` fails the same way. One + dev metapackage supplies all of them, and touches no source: + + ```bash + # Assumes AMD's Developer Cloud image, whose host already carries this repo's + # signing key at /etc/apt/keyrings/amdrocm.gpg -- copy it into the container. + # Anywhere else, install AMD's key under that path first: repo.amd.com serves + # none over HTTP, and the repo.radeon.com key signs a different repository, so + # apt-get update fails with NO_PUBKEY without it. + echo 'deb [arch=amd64 signed-by=/etc/apt/keyrings/amdrocm.gpg] https://repo.amd.com/rocm/packages-multi-arch/ubuntu2404 stable main' \ + > /etc/apt/sources.list.d/rocm.list + apt-get update && apt-get install -y amdrocm-core-dev7.14 + export CPLUS_INCLUDE_PATH=/opt/rocm/core-7.14/include + export LIBRARY_PATH=/opt/rocm/core-7.14/lib + ``` + + `LIBRARY_PATH` is separate from the headers: the image ships + `libamdhip64.so.7` with no development symlink, so the final link cannot + resolve `-lamdhip64` without it. Two near misses are worth naming, because + both fail far from their cause. Do not append `/usr/include` to + `CPLUS_INCLUDE_PATH` — searching it ahead of the compiler's own directories + defeats libstdc++'s `#include_next`, and the build dies on `stdlib.h: No such + file or directory` before it reaches a single ROCm header. And do not reach + for Ubuntu's `librocthrust-dev` / `librocprim-dev`: they are ROCm 5.7, and + they install a 5.7 HIP into `/usr/include` that shadows the image's 7.14 + headers, producing a wall of errors inside `amd_warp_sync_functions.h`. + +### Portability + +The kernels use no shared memory, no `__syncthreads()` and no warp-level +primitives. Nothing depends on the warp or wavefront width, so CDNA's 64-wide +wavefront against NVIDIA's 32-wide warp affects occupancy and nothing else. The +three block widths in `best_block_dim()` are multiples of 64, so neither +platform schedules a partial wave. + +Those thresholds were tuned on NVIDIA. What was measured on CDNA is the width +the heuristic actually picks at swin-tiny's `C` — 64 at both stages, since 96 and +192 are below the first threshold — by rebuilding with `-DSWIN_WP_BLOCK_DIM=N` +(stage 1 / stage 2, float32, MI300X, fused time): + +| block width | stage 1 | stage 2 | +|---|---:|---:| +| 64 (what the heuristic picks here) | **0.598 ms** | **0.188 ms** | +| 256 | 0.632 ms | 0.212 ms | +| 1024 | 1.230 ms | 0.326 ms | + +The NVIDIA-tuned choice wins on CDNA too, and widening hurts monotonically — at +1024 the fused path becomes *slower than eager* (0.74×). Both configurations sit +in the same branch, so the 384 and 1024 thresholds themselves remain untested on +CDNA. The grid is fixed by +the window geometry, so threads past `C` have nothing to do; swin-tiny's `C` is +96 at stage 1, and a 1024-wide block leaves 928 lanes idle. + +The one substantive difference between the two platforms is the launch stream. +Upstream passes `0`, the null stream; this version passes +`at::cuda::getCurrentCUDAStream()`. Some hipify versions rewrite that to the HIP +spelling and some leave it, which resolves to the HIP runtime anyway — both +reach the current stream. torch 2.8 and 2.10 rewrite it; 2.12 and 2.13 leave it, +and 2.12 is the version the MI300X build used, so the kept spelling is the one +that compiled under hipcc and ran bit-exact above. `hipify_check.py` still makes +no prediction: it reports the one that happened rather than requiring either, +prints the full symbol mapping, and fails if anything is left untranslated. diff --git a/kernels/window_process/benchmark.py b/kernels/window_process/benchmark.py new file mode 100644 index 000000000..eaa82f18f --- /dev/null +++ b/kernels/window_process/benchmark.py @@ -0,0 +1,162 @@ +# -------------------------------------------------------- +# Fused kernel for window process for SwinTransformer +# Copyright (c) 2022 Nvidia +# Licensed under The MIT License [see LICENSE for details] +# -------------------------------------------------------- + +"""Fused kernels against the PyTorch ops they replace, in time and peak memory. + +The eager path materialises the tensor twice -- torch.roll writes a copy, then +window_partition permutes and calls .contiguous() for another -- where the fused +kernel writes once. The gain is what that saved traffic is worth on the device, +which varies enough between vendors to be worth measuring rather than +predicting: README.md records what these configurations gave on an RTX 3080 and +on an MI300X. + + python benchmark.py # forward + backward, all dtypes + python benchmark.py --iters 200 --batch 64 +""" + +import argparse + +import torch + +import reference as ref + +try: + from window_process import WindowProcess, WindowProcessReverse +except ImportError as exc: # extension not built here + WindowProcess = WindowProcessReverse = None + IMPORT_ERROR = exc +else: + IMPORT_ERROR = None + + +# (name, B, H, W, C, window_h, window_w) +CONFIGS = [ + ('swin-tiny stage 1 56x56 w7', 192, 56, 56, 96, 7, 7), + ('swin-tiny stage 2 28x28 w7', 192, 28, 28, 192, 7, 7), + ('non-square 32x16 w8', 192, 32, 16, 96, 8, 8), + ('non-square window 16x64 w4x16', 192, 16, 64, 96, 4, 16), +] + + +def timed(fn, iters, warmup=10): + """Wall time per iteration in ms, and peak allocated memory in MiB.""" + for _ in range(warmup): + fn() + torch.cuda.synchronize() + torch.cuda.reset_peak_memory_stats() + + start, end = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True) + start.record() + for _ in range(iters): + fn() + end.record() + torch.cuda.synchronize() + + return start.elapsed_time(end) / iters, torch.cuda.max_memory_allocated() / 2 ** 20 + + +def make_runners(B, H, W, C, window_h, window_w, shift_h, shift_w, dtype, backward): + window = (window_h, window_w) + n = B * (H // window_h) * (W // window_w) + + x = torch.randn((B, H, W, C), dtype=dtype, device='cuda', requires_grad=backward) + grad = torch.randn((n, window_h, window_w, C), dtype=dtype, device='cuda') + + def eager(): + shifted = torch.roll(x, shifts=(-shift_h, -shift_w), dims=(1, 2)) + out = ref.window_partition(shifted, window) + if backward: + x.grad = None + out.backward(grad) + + def fused(): + out = WindowProcess.apply(x, B, H, W, C, (-shift_h, -shift_w), window) + if backward: + x.grad = None + out.backward(grad) + + return eager, fused + + +def make_reverse_runners(B, H, W, C, window_h, window_w, shift_h, shift_w, dtype, backward): + window = (window_h, window_w) + n = B * (H // window_h) * (W // window_w) + + x = torch.randn((n, window_h, window_w, C), dtype=dtype, device='cuda', + requires_grad=backward) + grad = torch.randn((B, H, W, C), dtype=dtype, device='cuda') + + def eager(): + merged = ref.window_reverse(x, window, H, W) + out = torch.roll(merged, shifts=(shift_h, shift_w), dims=(1, 2)) + if backward: + x.grad = None + out.backward(grad) + + def fused(): + out = WindowProcessReverse.apply(x, B, H, W, C, (shift_h, shift_w), window) + if backward: + x.grad = None + out.backward(grad) + + return eager, fused + + +def run(args): + dtypes = [torch.float32, torch.float16] + if torch.cuda.is_bf16_supported(): + dtypes.append(torch.bfloat16) + + print(f'device: {torch.cuda.get_device_name()}') + print(f'torch: {torch.__version__} (cuda {torch.version.cuda}, hip {torch.version.hip})') + print(f'mode: {"forward + backward" if not args.forward_only else "forward only"}, ' + f'{args.iters} iterations after 10 warmup\n') + + header = (f'| {"config":<32} | {"op":<7} | {"dtype":<9} | {"eager ms":>9} | ' + f'{"fused ms":>9} | {"speedup":>8} | {"eager MiB":>10} | {"fused MiB":>10} |') + print(header) + print('|' + '|'.join('-' * (len(c) + 2) for c in header.split('|')[1:-1]) + '|') + + for name, B, H, W, C, window_h, window_w in CONFIGS: + B = args.batch or B + shift_h, shift_w = window_h // 2, window_w // 2 + for op, factory in (('partition', make_runners), ('merge', make_reverse_runners)): + for dtype in dtypes: + eager, fused = factory(B, H, W, C, window_h, window_w, + shift_h, shift_w, dtype, + backward=not args.forward_only) + e_ms, e_mem = timed(eager, args.iters) + f_ms, f_mem = timed(fused, args.iters) + print(f'| {name:<32} | {op:<7} | {str(dtype).replace("torch.", ""):<9} | ' + f'{e_ms:>9.3f} | {f_ms:>9.3f} | {e_ms / f_ms:>7.2f}x | ' + f'{e_mem:>10.1f} | {f_mem:>10.1f} |') + del eager, fused + torch.cuda.empty_cache() + + +def main(): + parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) + parser.add_argument('--iters', type=int, default=100, + help='timed iterations per configuration') + parser.add_argument('--batch', type=int, default=None, + help='override the batch size of every config') + parser.add_argument('--forward-only', action='store_true', + help='skip the backward pass') + args = parser.parse_args() + + if not torch.cuda.is_available(): + raise SystemExit('benchmark.py requires a CUDA device') + if IMPORT_ERROR is not None: + raise SystemExit( + f'the swin_window_process extension is not importable ({IMPORT_ERROR}); ' + 'build it with: python setup.py install') + + torch.manual_seed(0) + run(args) + + +if __name__ == '__main__': + main() diff --git a/kernels/window_process/hipify_check.py b/kernels/window_process/hipify_check.py new file mode 100644 index 000000000..f54fcd22d --- /dev/null +++ b/kernels/window_process/hipify_check.py @@ -0,0 +1,237 @@ +# -------------------------------------------------------- +# Fused kernel for window process for SwinTransformer +# Copyright (c) 2022 Nvidia +# Licensed under The MIT License [see LICENSE for details] +# -------------------------------------------------------- + +"""Check that the CUDA sources translate cleanly to HIP for ROCm. + +AMD GPUs do not compile CUDA. They compile HIP, a nearly identical language, and +PyTorch gets from one to the other with a source-to-source translator called +hipify: CUDA names go in, HIP names come out. CUDAExtension runs it on its own +when torch.version.hip is set, so a ROCm build needs no second copy of these +sources -- but it gives no signal about whether the translation came out +complete. This is that signal, and because hipify is pure Python it can be had +with no AMD GPU (with no GPU at all) on a CPU runner in CI. + +Three things fail the check: a CUDA-only header that survives translation, a +name that is already portable but gets rewritten anyway, and a kernel launch +that reaches HIP without naming a stream. Symbols that translate, and the two +PyTorch spellings that are valid either way, are reported rather than judged. + + python hipify_check.py # verify + python hipify_check.py --diff # verify, and print the generated HIP +""" + +import argparse +import difflib +import os +import re +import shutil +import sys +import tempfile + +from torch.utils.hipify import hipify_python + + +SOURCES = ['swin_window_process.cpp', 'swin_window_process_kernel.cu'] + +# CUDA header files, which simply do not exist on an AMD machine. hipify has to +# rewrite every one of them to its HIP counterpart -- cuda_runtime.h becomes +# hip/hip_runtime.h, and so on. If one of these names is still there afterwards, +# the ROCm build stops at "no such file", so finding one is a failure. +MUST_BE_TRANSLATED = [ + 'cuda.h', + 'cuda_runtime.h', + 'cuda_fp16.h', + 'ATen/cuda/CUDAContext.h', + 'c10/cuda/CUDAException.h', +] + +# PyTorch functions and macros that work on ROCm under either name. Some hipify +# versions rename them to an at::hip / C10_HIP spelling and some leave the CUDA +# spelling in place, which resolves to the HIP runtime anyway: torch 2.8 and 2.10 +# rewrite them, 2.12 and 2.13 do not. Both spellings reach the current HIP +# stream -- the kept form is the one the MI300X build below actually compiled and +# ran, on torch 2.12 -- so the check makes no prediction and accepts either, +# reporting which one happened. That is why it passes unchanged on CI legs that +# pin different torch versions. +CUDA_COMPAT_SHIMS = [ + 'at::cuda::getCurrentCUDAStream', + 'C10_CUDA_KERNEL_LAUNCH_CHECK', +] + +# Names that mean the same thing in both languages: HIP spells blockIdx, +# __global__ and the rest exactly as CUDA does, so hipify is meant to leave them +# alone. One coming back renamed would mean the translation went wrong, so that +# is a failure too. +KNOWN_PORTABLE = [ + '__global__', + 'blockIdx', + 'threadIdx', + 'blockDim', + 'gridDim', + 'dim3', + 'AT_DISPATCH_FLOATING_TYPES_AND2', + 'at::ScalarType::BFloat16', +] + +# __ldg is deliberately not in the list above. It is portable as a *symbol* -- +# HIP defines it, and hipify leaves it alone -- but not for every type, which is +# the whole point of the caveat printed on success. Since the kernels read +# through SWIN_WP_LDG, the only __ldg left in the source is inside the #else +# branch of that macro, which hipcc never compiles: asserting that the token +# survives translation would be asserting something about dead code. The macro +# is reported instead, so the reader sees which arm the build will take. +LDG_MACRO = 'SWIN_WP_LDG' + +# Printed on success, because passing this check proves less than it sounds like. +SYMBOL_LEVEL_CAVEAT = """\ +Symbol level only: a symbol can translate and still not compile, because HIP and +CUDA do not always give it the same overload set. __ldg is the case in point -- +it has no HIP overload for c10::Half. Only a real build finds that.""" + + +def read(path): + with open(path) as handle: + return handle.read() + + +def hipify_into(destination): + """Translate SOURCES out of place and return {source: hipified path}.""" + here = os.path.dirname(os.path.abspath(__file__)) + staging = os.path.join(destination, 'src') + os.makedirs(staging) + for name in SOURCES: + shutil.copy(os.path.join(here, name), staging) + + results = hipify_python.hipify( + project_directory=staging, + output_directory=os.path.join(destination, 'out'), + includes=('*',), + is_pytorch_extension=True, + out_of_place_only=True, + show_progress=False, + ) + + mapping = {} + for source, result in results.items(): + hipified = getattr(result, 'hipified_path', None) or source + mapping[os.path.basename(source)] = hipified + return staging, mapping + + +def report(staging, mapping, show_diff): + """Classify each source's device-specific symbols and return the failures.""" + failures = [] + + for name in SOURCES: + original = os.path.join(staging, name) + hipified = mapping.get(name) + if hipified is None or not os.path.exists(hipified): + failures.append(f'{name}: hipify produced no output') + continue + + before = read(original) + after = read(hipified) + + print(f'\n{name} -> {os.path.basename(hipified)}') + + translated, survived = [], [] + for token in MUST_BE_TRANSLATED: + if token not in before: + continue + (survived if token in after else translated).append(token) + + for token in translated: + print(f' translated {token}') + for token in survived: + print(f' SURVIVED {token}') + failures.append(f'{name}: {token} was not translated') + + shims = [t for t in CUDA_COMPAT_SHIMS if t in before] + for token in shims: + how = 'rewritten' if token not in after else 'kept' + print(f' shim {token} ({how})') + + if LDG_MACRO in before: + print(f' macro {LDG_MACRO} ' + '(__ldg on NVIDIA, plain load on AMD -- see the note below)') + + portable = [t for t in KNOWN_PORTABLE if t in before] + for token in portable: + if token in after: + print(f' portable {token} (native in HIP, left as is)') + else: + failures.append(f'{name}: {token} was rewritten but should be portable') + + if not translated and not survived and not shims and not portable \ + and LDG_MACRO not in before: + print(' no symbol from the watch lists appears in this file') + + if show_diff: + diff = difflib.unified_diff( + before.splitlines(True), after.splitlines(True), + fromfile=name, tofile=os.path.basename(hipified)) + sys.stdout.writelines(diff) + + return failures + + +def check_launch_form(mapping): + """Every kernel launch must still carry an explicit stream after translation. + + A launch that reaches HIP on the null stream is the ROCm form of the + default-stream bug. + + The spelling of a launch is not fixed, so this accepts all of it: hipify may + turn the triple chevron into hipLaunchKernelGGL or leave it alone, and may + wrap the call over several lines. What it requires is narrow and textual: the + launch must name a current-stream getter under either spelling. A launch + handed a stream by any other route -- a variable, getStreamFromPool() -- is + rejected too, which is strict rather than wrong for this source. + """ + hipified = mapping.get('swin_window_process_kernel.cu') + if hipified is None: + return ['no hipified kernel source to inspect'] + + text = read(hipified) + calls = re.findall(r'hipLaunchKernelGGL\s*\((.*?)\)\s*;', text, re.DOTALL) + calls += re.findall(r'<<<(.*?)>>>', text, re.DOTALL) + if not calls: + return ['no kernel launch found in the translated source'] + + stream_getters = ('getCurrentHIPStream', 'getCurrentCUDAStream') + failures = [] + for call in calls: + if not any(g in call for g in stream_getters): + failures.append('launch without an explicit stream: ' + + ' '.join(call.split())[:80]) + if not failures: + print(f'\n{len(calls)} kernel launches, all carrying an explicit stream') + return failures + + +def main(): + parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) + parser.add_argument('--diff', action='store_true', + help='print the generated HIP source as a unified diff') + args = parser.parse_args() + + with tempfile.TemporaryDirectory() as tmp: + staging, mapping = hipify_into(tmp) + failures = report(staging, mapping, args.diff) + failures += check_launch_form(mapping) + + if failures: + print('\nFAILED') + for failure in failures: + print(f' {failure}') + raise SystemExit(1) + + print('\nOK: the CUDA sources translate to HIP with no unmapped symbols.') + print(SYMBOL_LEVEL_CAVEAT) + + +if __name__ == '__main__': + main() diff --git a/kernels/window_process/reference.py b/kernels/window_process/reference.py new file mode 100644 index 000000000..215e0426d --- /dev/null +++ b/kernels/window_process/reference.py @@ -0,0 +1,226 @@ +# -------------------------------------------------------- +# Fused kernel for window process for SwinTransformer +# Copyright (c) 2022 Nvidia +# Licensed under The MIT License [see LICENSE for details] +# -------------------------------------------------------- +# Device-independent transcription of the index arithmetic used by the four +# CUDA kernels in swin_window_process_kernel.cu. +# +# The kernels perform no arithmetic on the tensor values: every one is a pure +# gather, and all of their logic lives in computing `input_offset` from +# `blockIdx`. Transcribing that computation to vectorised PyTorch therefore +# reproduces the kernels exactly -- which is what makes them testable with no +# compiled extension and no GPU, on a CI runner (see test_index_math.py). +# +# Exactly, but not unconditionally. Python's % and // floor, C++'s truncate, and +# in the kernels the expression is unsigned as well, so it wraps modulo 2^32 +# where this file would carry a negative value through. The two agree for as +# long as the operand of every % stays non-negative, which the `+ H` and `+ W` +# terms guarantee while |shift| <= H and |shift| <= W -- a wider range than the +# launcher's TORCH_CHECKs admit, so inside the supported domain the two are the +# same arithmetic. Outside it they are not, and this file is the more forgiving +# of the two. +# -------------------------------------------------------- + +import torch + + +def to_pair(value): + """Accept an int (isotropic) or a (h, w) iterable and return a (h, w) tuple.""" + if isinstance(value, int): + return value, value + h, w = value + return int(h), int(w) + + +def window_partition(x, window_size): + """(B, H, W, C) -> (B * nH * nW, window_h, window_w, C).""" + window_h, window_w = to_pair(window_size) + B, H, W, C = x.shape + x = x.view(B, H // window_h, window_h, W // window_w, window_w, C) + return x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, window_h, window_w, C) + + +def window_reverse(windows, window_size, H, W): + """(B * nH * nW, window_h, window_w, C) -> (B, H, W, C).""" + window_h, window_w = to_pair(window_size) + C = windows.shape[-1] + B = windows.shape[0] // ((H // window_h) * (W // window_w)) + x = windows.view(B, H // window_h, W // window_w, window_h, window_w, C) + return x.permute(0, 1, 3, 2, 4, 5).contiguous().view(B, H, W, C) + + +# -------------------------------------------------------------------------- +# Index maps +# +# Each function returns the flat index, into the *spatial* grid (B * H * W) or +# into the *window* grid (B * nH * nW * window_h * window_w), that every output +# element reads from. Channels are omitted: the kernels loop over C with a fixed +# stride of 1, so the channel dimension contributes an identical offset to every +# element and can be factored out by gathering over rows of a (-1, C) view. +# +# `shift_h` / `shift_w` follow the kernel convention, which is the *negated* +# torch.roll shift on the partition path — callers pass -shift_size to the +# forward kernels and +shift_size to the reverse ones. See the identities +# asserted in test_index_math.py. +# -------------------------------------------------------------------------- + + +def roll_and_window_partition_forward_index(B, H, W, shift_h, shift_w, window_h, window_w): + """K1: roll + window partition. Reads from the spatial grid. + + Returns: + (B * nH * nW, window_h, window_w) int64 tensor of flat indices into + a (B * H * W, C) view of the input. + """ + nH, nW = H // window_h, W // window_w + flat_win = torch.arange(B * nH * nW).view(-1, 1, 1) # blockIdx.z + ly = torch.arange(window_h).view(1, -1, 1) # blockIdx.y + lx = torch.arange(window_w).view(1, 1, -1) # blockIdx.x + + b = flat_win // (nH * nW) + row = (flat_win % (nH * nW) // nW * window_h + ly - shift_h + H) % H + col = (flat_win % nW * window_w + lx - shift_w + W) % W + return b * (H * W) + row * W + col + + +def roll_and_window_partition_backward_index(B, H, W, shift_h, shift_w, window_h, window_w): + """K2: backward of K1. Reads from the window grid. + + Returns: + (B, H, W) int64 tensor of flat indices into a + (B * nH * nW * window_h * window_w, C) view of the incoming gradient. + """ + nH, nW = H // window_h, W // window_w + b = torch.arange(B).view(-1, 1, 1) # blockIdx.z + y = torch.arange(H).view(1, -1, 1) # blockIdx.y + x = torch.arange(W).view(1, 1, -1) # blockIdx.x + + src_y = (y + shift_h + H) % H + src_x = (x + shift_w + W) % W + win = b * (nH * nW) + src_y // window_h * nW + src_x // window_w + return win * (window_h * window_w) + (src_y % window_h) * window_w + (src_x % window_w) + + +def window_merge_and_roll_forward_index( + B, H, W, shift_h, shift_w, window_h, window_w, + *, + legacy_row_stride=False, + legacy_intra_modulo=False, +): + """K3: window merge + reverse roll. Reads from the window grid. + + Args: + legacy_row_stride: reproduce the upstream ``* nH`` term. The stride + between consecutive window *rows* is nW (there are nW windows per + row in the row-major layout ``b * nH * nW + wrow * nW + wcol``), so + ``* nH`` is only correct when nH == nW, i.e. for square feature + maps. This is the bug fixed by this branch. + legacy_intra_modulo: reproduce the upstream intra-window modulo, which + omits the ``% H`` / ``% W`` reduction before ``% window_h`` / + ``% window_w``. This is equivalent to the explicit form whenever + H % window_h == 0 and W % window_w == 0, which the launcher already + requires; the flag exists to make that equivalence testable rather + than assumed. + + Returns: + (B, H, W) int64 tensor of flat indices into a + (B * nH * nW * window_h * window_w, C) view of the input. + """ + nH, nW = H // window_h, W // window_w + b = torch.arange(B).view(-1, 1, 1) # blockIdx.z + y = torch.arange(H).view(1, -1, 1) # blockIdx.y + x = torch.arange(W).view(1, 1, -1) # blockIdx.x + + src_y = (y - shift_h + H) % H + src_x = (x - shift_w + W) % W + + row_stride = nH if legacy_row_stride else nW + win = b * (nH * nW) + src_y // window_h * row_stride + src_x // window_w + + if legacy_intra_modulo: + intra_y = (y - shift_h + H) % window_h + intra_x = (x - shift_w + W) % window_w + else: + intra_y = src_y % window_h + intra_x = src_x % window_w + + return win * (window_h * window_w) + intra_y * window_w + intra_x + + +def window_merge_and_roll_backward_index(B, H, W, shift_h, shift_w, window_h, window_w): + """K4: backward of K3. Reads from the spatial grid. + + Returns: + (B * nH * nW, window_h, window_w) int64 tensor of flat indices into + a (B * H * W, C) view of the incoming gradient. + """ + nH, nW = H // window_h, W // window_w + flat_win = torch.arange(B * nH * nW).view(-1, 1, 1) # blockIdx.z + ly = torch.arange(window_h).view(1, -1, 1) # blockIdx.y + lx = torch.arange(window_w).view(1, 1, -1) # blockIdx.x + + b = flat_win // (nH * nW) + row = (flat_win % (nH * nW) // nW * window_h + ly + shift_h + H) % H + col = (flat_win % nW * window_w + lx + shift_w + W) % W + return b * (H * W) + row * W + col + + +# -------------------------------------------------------------------------- +# Gathers +# -------------------------------------------------------------------------- + + +def _gather(src, index, out_shape): + """Apply a flat index map to the leading dimensions of ``src``. + + Args: + src: tensor whose last dimension is C. + index: int64 tensor of flat indices into ``src.view(-1, C)``. + out_shape: spatial shape of the result, C appended by this function. + """ + C = src.shape[-1] + flat = src.reshape(-1, C) + return flat[index.reshape(-1)].view(*out_shape, C) + + +def roll_and_window_partition_forward(x, shift_size, window_size): + """Reference implementation of K1. ``x``: (B, H, W, C).""" + shift_h, shift_w = to_pair(shift_size) + window_h, window_w = to_pair(window_size) + B, H, W, _ = x.shape + index = roll_and_window_partition_forward_index( + B, H, W, shift_h, shift_w, window_h, window_w) + nH, nW = H // window_h, W // window_w + return _gather(x, index, (B * nH * nW, window_h, window_w)) + + +def roll_and_window_partition_backward(grad_in, shift_size, window_size, H, W): + """Reference implementation of K2. ``grad_in``: (B*nH*nW, window_h, window_w, C).""" + shift_h, shift_w = to_pair(shift_size) + window_h, window_w = to_pair(window_size) + B = grad_in.shape[0] // ((H // window_h) * (W // window_w)) + index = roll_and_window_partition_backward_index( + B, H, W, shift_h, shift_w, window_h, window_w) + return _gather(grad_in, index, (B, H, W)) + + +def window_merge_and_roll_forward(x, shift_size, window_size, H, W, **legacy): + """Reference implementation of K3. ``x``: (B*nH*nW, window_h, window_w, C).""" + shift_h, shift_w = to_pair(shift_size) + window_h, window_w = to_pair(window_size) + B = x.shape[0] // ((H // window_h) * (W // window_w)) + index = window_merge_and_roll_forward_index( + B, H, W, shift_h, shift_w, window_h, window_w, **legacy) + return _gather(x, index, (B, H, W)) + + +def window_merge_and_roll_backward(grad_in, shift_size, window_size): + """Reference implementation of K4. ``grad_in``: (B, H, W, C).""" + shift_h, shift_w = to_pair(shift_size) + window_h, window_w = to_pair(window_size) + B, H, W, _ = grad_in.shape + index = window_merge_and_roll_backward_index( + B, H, W, shift_h, shift_w, window_h, window_w) + nH, nW = H // window_h, W // window_w + return _gather(grad_in, index, (B * nH * nW, window_h, window_w)) diff --git a/kernels/window_process/swin_window_process.cpp b/kernels/window_process/swin_window_process.cpp index a7f15b048..fec0bc8b6 100644 --- a/kernels/window_process/swin_window_process.cpp +++ b/kernels/window_process/swin_window_process.cpp @@ -17,6 +17,9 @@ #include #include +#include +#include + at::Tensor roll_and_window_partition_forward_cuda( at::Tensor & input, @@ -25,8 +28,10 @@ at::Tensor roll_and_window_partition_forward_cuda( const int H, const int W, const int C, - const int shift_size, - const int window_size); + const int shift_h, + const int shift_w, + const int window_h, + const int window_w); at::Tensor roll_and_window_partition_backward_cuda( @@ -36,8 +41,10 @@ at::Tensor roll_and_window_partition_backward_cuda( const int H, const int W, const int C, - const int shift_size, - const int window_size); + const int shift_h, + const int shift_w, + const int window_h, + const int window_w); at::Tensor window_merge_and_roll_forward_cuda( @@ -47,8 +54,10 @@ at::Tensor window_merge_and_roll_forward_cuda( const int H, const int W, const int C, - const int shift_size, - const int window_size); + const int shift_h, + const int shift_w, + const int window_h, + const int window_w); at::Tensor window_merge_and_roll_backward_cuda( at::Tensor & grad_in, @@ -57,15 +66,76 @@ at::Tensor window_merge_and_roll_backward_cuda( const int H, const int W, const int C, - const int shift_size, - const int window_size); + const int shift_h, + const int shift_w, + const int window_h, + const int window_w); -#define CHECK_CUDA(x) AT_ASSERTM(x.type().is_cuda(), #x " must be a CUDA tensor") +#define CHECK_CUDA(x) AT_ASSERTM(x.is_cuda(), #x " must be a CUDA tensor") #define CHECK_CONTIGUOUS(x) AT_ASSERTM(x.is_contiguous(), #x " must be contiguous") #define CHECK_INPUT(x) CHECK_CUDA(x); CHECK_CONTIGUOUS(x) +// The kernels derive nH = H / window_h and nW = W / window_w with integer +// division and index with int arithmetic. All four entry points share the same +// constraints -- the window layout and the spatial layout hold the same number +// of elements -- so one check covers them. +// +// The checks below are not all guarding the same failure. Divisibility, numel +// and the int32 bound catch what used to be a silent wrong read or an +// out-of-bounds one, with no error raised at any point. The grid bounds catch a +// launch that used to fail asynchronously somewhere else, leaving the output +// uninitialised. The shift bound is neither: the formula stays correct for +// window <= |shift| <= H (and <= W on the other axis), which is simply a larger +// roll. Past that the `+ H` no longer keeps the operand of `%` non-negative, +// and since the block index is unsigned the expression wraps modulo 2^32 rather +// than going negative -- `% H` still yields an in-range index, so the read is +// silently wrong rather than out of bounds. The bound is here because it is the +// contract the model already documents (`0 <= shift_size < window_size` in +// models/swin_transformer.py), made explicit at the boundary. +static void check_window_args( + const at::Tensor & tensor, + const int B, + const int H, + const int W, + const int C, + const int shift_h, + const int shift_w, + const int window_h, + const int window_w){ + + TORCH_CHECK(B > 0 && H > 0 && W > 0 && C > 0, + "B, H, W and C must be positive, got (", B, ", ", H, ", ", W, ", ", C, ")"); + TORCH_CHECK(window_h > 0 && window_w > 0, + "window size must be positive, got (", window_h, ", ", window_w, ")"); + TORCH_CHECK(H % window_h == 0, + "H (", H, ") must be divisible by window_h (", window_h, ")"); + TORCH_CHECK(W % window_w == 0, + "W (", W, ") must be divisible by window_w (", window_w, ")"); + TORCH_CHECK(std::abs(shift_h) < window_h && std::abs(shift_w) < window_w, + "|shift| must be smaller than the window size on each axis, got shift (", + shift_h, ", ", shift_w, ") for window (", window_h, ", ", window_w, ")"); + + const int64_t numel = static_cast(B) * H * W * C; + TORCH_CHECK(tensor.numel() == numel, + "expected ", numel, " elements for B=", B, ", H=", H, ", W=", W, ", C=", C, + ", got ", tensor.numel()); + // Offsets are computed in int inside the kernels, as they are upstream. + TORCH_CHECK(numel <= static_cast(std::numeric_limits::max()), + "tensor holds ", numel, " elements; kernel offsets are computed in int32 " + "and would overflow"); + + // grid.z is B * nH * nW for the partition kernels and B for the merge ones; + // grid.y is window_h or H. Both are capped at 65535 by the CUDA driver. + const int64_t num_windows = static_cast(B) * (H / window_h) * (W / window_w); + TORCH_CHECK(num_windows <= 65535, + "B * nH * nW = ", num_windows, " exceeds the CUDA grid.z limit of 65535"); + TORCH_CHECK(H <= 65535, + "H = ", H, " exceeds the CUDA grid.y limit of 65535"); +} + + at::Tensor roll_and_window_partition_forward( at::Tensor & input, @@ -74,10 +144,13 @@ at::Tensor roll_and_window_partition_forward( const int H, const int W, const int C, - const int shift_size, - const int window_size){ + const int shift_h, + const int shift_w, + const int window_h, + const int window_w){ CHECK_INPUT(input); - return roll_and_window_partition_forward_cuda(input, B, H, W, C, shift_size, window_size); + check_window_args(input, B, H, W, C, shift_h, shift_w, window_h, window_w); + return roll_and_window_partition_forward_cuda(input, B, H, W, C, shift_h, shift_w, window_h, window_w); } @@ -88,10 +161,13 @@ at::Tensor roll_and_window_partition_backward( const int H, const int W, const int C, - const int shift_size, - const int window_size){ + const int shift_h, + const int shift_w, + const int window_h, + const int window_w){ CHECK_INPUT(grad_in); - return roll_and_window_partition_backward_cuda(grad_in, B, H, W, C, shift_size, window_size); + check_window_args(grad_in, B, H, W, C, shift_h, shift_w, window_h, window_w); + return roll_and_window_partition_backward_cuda(grad_in, B, H, W, C, shift_h, shift_w, window_h, window_w); } @@ -102,10 +178,13 @@ at::Tensor window_merge_and_roll_forward( const int H, const int W, const int C, - const int shift_size, - const int window_size){ + const int shift_h, + const int shift_w, + const int window_h, + const int window_w){ CHECK_INPUT(input); - return window_merge_and_roll_forward_cuda(input, B, H, W, C, shift_size, window_size); + check_window_args(input, B, H, W, C, shift_h, shift_w, window_h, window_w); + return window_merge_and_roll_forward_cuda(input, B, H, W, C, shift_h, shift_w, window_h, window_w); } @@ -116,10 +195,13 @@ at::Tensor window_merge_and_roll_backward( const int H, const int W, const int C, - const int shift_size, - const int window_size){ + const int shift_h, + const int shift_w, + const int window_h, + const int window_w){ CHECK_INPUT(grad_in); - return window_merge_and_roll_backward_cuda(grad_in, B, H, W, C, shift_size, window_size); + check_window_args(grad_in, B, H, W, C, shift_h, shift_w, window_h, window_w); + return window_merge_and_roll_backward_cuda(grad_in, B, H, W, C, shift_h, shift_w, window_h, window_w); } diff --git a/kernels/window_process/swin_window_process_kernel.cu b/kernels/window_process/swin_window_process_kernel.cu index 4976949ee..f50e667c1 100644 --- a/kernels/window_process/swin_window_process_kernel.cu +++ b/kernels/window_process/swin_window_process_kernel.cu @@ -15,13 +15,49 @@ */ #include +#include +#include +#include #include #include #include -#include #include +// Read-only cached load: __ldg on NVIDIA, a plain load on AMD. +// +// __ldg covers every dispatched type only under CUDA. c10::Half and +// c10::BFloat16 come from torch/headeronly/util and are not guarded alike: +// BFloat16 sits behind `__CUDACC__ || __HIPCC__` and degrades to *ptr on ROCm, +// Half behind `(__CUDA_ARCH__ >= 350) || (__clang__ && __CUDA__)`, which no +// hipcc build satisfies. Half is in the dispatch, so this file does not compile +// under ROCm as written. The hint is advisory in both languages -- it asks for +// the read-only path and may be ignored -- so dropping it on AMD changes no +// result. A macro rather than an inline function so that the NVIDIA build still +// sees the same __ldg(ptr) tokens as before. +#if defined(__HIP_PLATFORM_AMD__) +#define SWIN_WP_LDG(ptr) (*(ptr)) +#else +#define SWIN_WP_LDG(ptr) __ldg(ptr) +#endif + +// Width of the block that walks the channel dimension. +// +// No shared memory, no __syncthreads(), no warp-level primitives: this width +// affects occupancy only, never correctness, and nothing here depends on the +// warp or wavefront width -- CDNA's 64 against NVIDIA's 32 is immaterial. All +// three values below are multiples of 64, so neither platform runs a partial +// wave. +// +// The thresholds are NVIDIA-tuned. What CDNA confirms is the width they pick at +// swin-tiny's C, 64 at both stages: on an MI300X 64 beats 256 beats 1024, and +// at 1024 the fused path loses to eager. Both stages land in the first branch, +// so the 384 and 1024 thresholds themselves remain untested there. Override +// with -DSWIN_WP_BLOCK_DIM=N. int best_block_dim(int feat_dim){ +#ifdef SWIN_WP_BLOCK_DIM + (void)feat_dim; + return SWIN_WP_BLOCK_DIM; +#else int best_dim; if (feat_dim < 384){ best_dim = 64; @@ -35,19 +71,31 @@ int best_block_dim(int feat_dim){ } } return best_dim; +#endif } +// The four kernels below are pure gathers: each one *is* the input_offset it +// computes. reference.py transcribes those expressions to PyTorch so they can +// be tested on a CPU. Nothing binds the two files -- only a build checks this +// source -- so if you change the arithmetic here, change reference.py too, or +// test_index_math.py will go on passing against the old formula. + +// input: [B, H, W, C] +// output: [B*nH*nW, window_h, window_w, C] +// grid: (window_w, window_h, B*nH*nW) -- x indexes columns, y indexes rows template __global__ void roll_and_window_partition_forward_cuda_kernel( - T* input, - T* output, + T* input, + T* output, const int B, const int H, const int W, const int C, - const int shift_size, - const int window_size, + const int shift_h, + const int shift_w, + const int window_h, + const int window_w, const int nH, const int nW){ // start @@ -57,24 +105,29 @@ __global__ void roll_and_window_partition_forward_cuda_kernel( for (int i = index; i < C; i += blockDim.x) { offset = ((blockIdx.z * gridDim.y + blockIdx.y) * gridDim.x + blockIdx.x) * C + i; // C = blocksize int input_offset = blockIdx.z / (nH * nW) * H * W * C + - (blockIdx.z % (nH * nW) / nW * window_size + blockIdx.y - shift_size + H) % H * W * C + - (blockIdx.z % nW * window_size + blockIdx.x - shift_size + W) % W * C + + (blockIdx.z % (nH * nW) / nW * window_h + blockIdx.y - shift_h + H) % H * W * C + + (blockIdx.z % nW * window_w + blockIdx.x - shift_w + W) % W * C + i; - output[offset] = (T)(__ldg(input + input_offset)); + output[offset] = (T)(SWIN_WP_LDG(input + input_offset)); } } +// grad_in: [B*nH*nW, window_h, window_w, C] +// grad_out: [B, H, W, C] +// grid: (W, H, B) template __global__ void roll_and_window_partition_backward_cuda_kernel( - T* grad_in, - T* grad_out, + T* grad_in, + T* grad_out, const int B, const int H, const int W, const int C, - const int shift_size, - const int window_size, + const int shift_h, + const int shift_w, + const int window_h, + const int window_w, const int nH, const int nW){ // start @@ -82,26 +135,31 @@ __global__ void roll_and_window_partition_backward_cuda_kernel( int offset; for (int i = index; i < C; i += blockDim.x) { offset = ((blockIdx.z * gridDim.y + blockIdx.y) * gridDim.x + blockIdx.x) * C + i; // C = blocksize - int input_offset = - (blockIdx.z * nH * nW + (blockIdx.y + shift_size + H) % H / window_size * nW + (blockIdx.x + shift_size + W) % W / window_size) * window_size * window_size * C + - (blockIdx.y + shift_size + H ) % H % window_size * window_size * C + - (blockIdx.x + shift_size + W ) % W % window_size * C + + int input_offset = + (blockIdx.z * nH * nW + (blockIdx.y + shift_h + H) % H / window_h * nW + (blockIdx.x + shift_w + W) % W / window_w) * window_h * window_w * C + + (blockIdx.y + shift_h + H ) % H % window_h * window_w * C + + (blockIdx.x + shift_w + W ) % W % window_w * C + i; - grad_out[offset] = (T)(__ldg(grad_in + input_offset)); + grad_out[offset] = (T)(SWIN_WP_LDG(grad_in + input_offset)); } } +// input: [B*nH*nW, window_h, window_w, C] +// output: [B, H, W, C] +// grid: (W, H, B) template __global__ void window_merge_and_roll_forward_cuda_kernel( - T* input, - T* output, + T* input, + T* output, const int B, const int H, const int W, const int C, - const int shift_size, - const int window_size, + const int shift_h, + const int shift_w, + const int window_h, + const int window_w, const int nH, const int nW){ // start @@ -109,27 +167,32 @@ __global__ void window_merge_and_roll_forward_cuda_kernel( int offset; for (int i = index; i < C; i += blockDim.x) { offset = ((blockIdx.z * gridDim.y + blockIdx.y) * gridDim.x + blockIdx.x) * C + i; // C = blocksize - int input_offset = - (blockIdx.z * nH * nW + (blockIdx.y - shift_size + H) % H / window_size * nH + (blockIdx.x - shift_size + W) % W / window_size) * window_size * window_size * C + - (blockIdx.y - shift_size + H) % window_size * window_size * C + - (blockIdx.x - shift_size + W) % window_size * C + + int input_offset = + (blockIdx.z * nH * nW + (blockIdx.y - shift_h + H) % H / window_h * nW + (blockIdx.x - shift_w + W) % W / window_w) * window_h * window_w * C + + (blockIdx.y - shift_h + H) % H % window_h * window_w * C + + (blockIdx.x - shift_w + W) % W % window_w * C + i; - output[offset] = (T)(__ldg(input + input_offset)); + output[offset] = (T)(SWIN_WP_LDG(input + input_offset)); } } +// grad_in: [B, H, W, C] +// grad_out: [B*nH*nW, window_h, window_w, C] +// grid: (window_w, window_h, B*nH*nW) template __global__ void window_merge_and_roll_backward_cuda_kernel( - T* grad_in, - T* grad_out, + T* grad_in, + T* grad_out, const int B, const int H, const int W, const int C, - const int shift_size, - const int window_size, + const int shift_h, + const int shift_w, + const int window_h, + const int window_w, const int nH, const int nW){ // start @@ -137,187 +200,188 @@ __global__ void window_merge_and_roll_backward_cuda_kernel( int offset; for (int i = index; i < C; i += blockDim.x) { offset = ((blockIdx.z * gridDim.y + blockIdx.y) * gridDim.x + blockIdx.x) * C + i; // C = blocksize - int input_offset = + int input_offset = (blockIdx.z / (nH * nW)) * H * W * C + - (blockIdx.z % (nH * nW) / nW * window_size + blockIdx.y + shift_size + H) % H * W * C + - (blockIdx.z % nW * window_size + blockIdx.x + shift_size + W) % W * C + + (blockIdx.z % (nH * nW) / nW * window_h + blockIdx.y + shift_h + H) % H * W * C + + (blockIdx.z % nW * window_w + blockIdx.x + shift_w + W) % W * C + i; - grad_out[offset] = (T)(__ldg(grad_in + input_offset)); + grad_out[offset] = (T)(SWIN_WP_LDG(grad_in + input_offset)); } } // input: [B, H, W, C] -// output: [B*nH*nW, window_size, window_size, C] +// output: [B*nH*nW, window_h, window_w, C] at::Tensor roll_and_window_partition_forward_cuda( - at::Tensor & input, + at::Tensor & input, //at::Tensor & output, const int B, const int H, const int W, const int C, - const int shift_size, - const int window_size){ - - int nH = H / window_size; - int nW = W / window_size; + const int shift_h, + const int shift_w, + const int window_h, + const int window_w){ + + int nH = H / window_h; + int nW = W / window_w; - dim3 grid(window_size, window_size, B * nH * nW); + dim3 grid(window_w, window_h, B * nH * nW); //dim3 block((C + 31) / 32 * 32); int blocknum = best_block_dim(C); dim3 block(blocknum); - at::Tensor output; - if (input.scalar_type() == torch::kFloat16){ - output = torch::empty({B*nH*nW, window_size, window_size, C}, torch::dtype(torch::kFloat16).device(torch::kCUDA).requires_grad(true)); - } - else{ - output = torch::empty({B*nH*nW, window_size, window_size, C}, torch::dtype(torch::kFloat32).device(torch::kCUDA).requires_grad(true)); - } + at::Tensor output = at::empty({B*nH*nW, window_h, window_w, C}, input.options()); - AT_DISPATCH_FLOATING_TYPES_AND_HALF(input.type(), "roll_and_window_partition_forward_cuda_kernel", ([&] { - roll_and_window_partition_forward_cuda_kernel<<>>( - input.data(), - output.data(), + AT_DISPATCH_FLOATING_TYPES_AND2(at::ScalarType::Half, at::ScalarType::BFloat16, + input.scalar_type(), "roll_and_window_partition_forward_cuda_kernel", ([&] { + roll_and_window_partition_forward_cuda_kernel<<>>( + input.data_ptr(), + output.data_ptr(), B, H, W, C, - shift_size, - window_size, + shift_h, + shift_w, + window_h, + window_w, nH, nW); })); + C10_CUDA_KERNEL_LAUNCH_CHECK(); return output; } -// grad_in: [B*nH*nW, window_size, window_size, C] +// grad_in: [B*nH*nW, window_h, window_w, C] // grad_out: [B, H, W, C] at::Tensor roll_and_window_partition_backward_cuda( - at::Tensor & grad_in, + at::Tensor & grad_in, const int B, const int H, const int W, const int C, - const int shift_size, - const int window_size){ - - int nH = H / window_size; - int nW = W / window_size; + const int shift_h, + const int shift_w, + const int window_h, + const int window_w){ + + int nH = H / window_h; + int nW = W / window_w; dim3 grid(W, H, B); //dim3 block((C + 31) / 32 * 32); int blocknum = best_block_dim(C); dim3 block(blocknum); - at::Tensor grad_out; - if (grad_in.scalar_type() == torch::kFloat16){ - grad_out = torch::empty({B, H, W, C}, torch::dtype(torch::kFloat16).device(torch::kCUDA).requires_grad(false)); - } - else{ - grad_out = torch::empty({B, H, W, C}, torch::dtype(torch::kFloat32).device(torch::kCUDA).requires_grad(false)); - } + at::Tensor grad_out = at::empty({B, H, W, C}, grad_in.options()); - AT_DISPATCH_FLOATING_TYPES_AND_HALF(grad_in.type(), "roll_and_window_partition_backward_cuda_kernel", ([&] { - roll_and_window_partition_backward_cuda_kernel<<>>( - grad_in.data(), - grad_out.data(), + AT_DISPATCH_FLOATING_TYPES_AND2(at::ScalarType::Half, at::ScalarType::BFloat16, + grad_in.scalar_type(), "roll_and_window_partition_backward_cuda_kernel", ([&] { + roll_and_window_partition_backward_cuda_kernel<<>>( + grad_in.data_ptr(), + grad_out.data_ptr(), B, H, W, C, - shift_size, - window_size, + shift_h, + shift_w, + window_h, + window_w, nH, nW); })); + C10_CUDA_KERNEL_LAUNCH_CHECK(); return grad_out; } -// input: [B*nH*nW, window_size, window_size, C] +// input: [B*nH*nW, window_h, window_w, C] // output: [B, H, W, C] at::Tensor window_merge_and_roll_forward_cuda( - at::Tensor & input, + at::Tensor & input, //at::Tensor & output, const int B, const int H, const int W, const int C, - const int shift_size, - const int window_size){ - - int nH = H / window_size; - int nW = W / window_size; + const int shift_h, + const int shift_w, + const int window_h, + const int window_w){ + + int nH = H / window_h; + int nW = W / window_w; dim3 grid(W, H, B); //dim3 block((C + 31) / 32 * 32); int blocknum = best_block_dim(C); dim3 block(blocknum); - //generate output tensor inside - at::Tensor output; - if (input.scalar_type() == torch::kFloat16){ - output = torch::empty({B, H, W, C}, torch::dtype(torch::kFloat16).device(torch::kCUDA).requires_grad(true)); - } - else{ - output = torch::empty({B, H, W, C}, torch::dtype(torch::kFloat32).device(torch::kCUDA).requires_grad(true)); - } + at::Tensor output = at::empty({B, H, W, C}, input.options()); - AT_DISPATCH_FLOATING_TYPES_AND_HALF(input.type(), "window_merge_and_roll_forward_cuda_kernel", ([&] { - window_merge_and_roll_forward_cuda_kernel<<>>( - input.data(), - output.data(), + AT_DISPATCH_FLOATING_TYPES_AND2(at::ScalarType::Half, at::ScalarType::BFloat16, + input.scalar_type(), "window_merge_and_roll_forward_cuda_kernel", ([&] { + window_merge_and_roll_forward_cuda_kernel<<>>( + input.data_ptr(), + output.data_ptr(), B, H, W, C, - shift_size, - window_size, + shift_h, + shift_w, + window_h, + window_w, nH, nW); })); + C10_CUDA_KERNEL_LAUNCH_CHECK(); return output; } +// grad_in: [B, H, W, C] +// grad_out: [B*nH*nW, window_h, window_w, C] at::Tensor window_merge_and_roll_backward_cuda( - at::Tensor & grad_in, + at::Tensor & grad_in, const int B, const int H, const int W, const int C, - const int shift_size, - const int window_size){ - - int nH = H / window_size; - int nW = W / window_size; + const int shift_h, + const int shift_w, + const int window_h, + const int window_w){ - dim3 grid(window_size, window_size, B * nH * nW); + int nH = H / window_h; + int nW = W / window_w; + + dim3 grid(window_w, window_h, B * nH * nW); //dim3 block((C + 31) / 32 * 32); int blocknum = best_block_dim(C); dim3 block(blocknum); - at::Tensor grad_out; - if (grad_in.scalar_type() == torch::kFloat16){ - grad_out = torch::empty({B*nH*nW, window_size, window_size, C}, torch::dtype(torch::kFloat16).device(torch::kCUDA).requires_grad(false)); - } - else{ - grad_out = torch::empty({B*nH*nW, window_size, window_size, C}, torch::dtype(torch::kFloat32).device(torch::kCUDA).requires_grad(false)); - } + at::Tensor grad_out = at::empty({B*nH*nW, window_h, window_w, C}, grad_in.options()); - AT_DISPATCH_FLOATING_TYPES_AND_HALF(grad_in.type(), "window_merge_and_roll_backward_cuda_kernel", ([&] { - window_merge_and_roll_backward_cuda_kernel<<>>( - grad_in.data(), - grad_out.data(), + AT_DISPATCH_FLOATING_TYPES_AND2(at::ScalarType::Half, at::ScalarType::BFloat16, + grad_in.scalar_type(), "window_merge_and_roll_backward_cuda_kernel", ([&] { + window_merge_and_roll_backward_cuda_kernel<<>>( + grad_in.data_ptr(), + grad_out.data_ptr(), B, H, W, C, - shift_size, - window_size, + shift_h, + shift_w, + window_h, + window_w, nH, nW); })); + C10_CUDA_KERNEL_LAUNCH_CHECK(); return grad_out; -} \ No newline at end of file +} diff --git a/kernels/window_process/test_index_math.py b/kernels/window_process/test_index_math.py new file mode 100644 index 000000000..22962b9a9 --- /dev/null +++ b/kernels/window_process/test_index_math.py @@ -0,0 +1,240 @@ +# -------------------------------------------------------- +# Fused kernel for window process for SwinTransformer +# Copyright (c) 2022 Nvidia +# Licensed under The MIT License [see LICENSE for details] +# -------------------------------------------------------- +# Correctness tests for the index arithmetic of the fused window kernels. +# +# These run on CPU against reference.py and need neither a GPU nor the compiled +# extension, so they can gate the index math in CI. The GPU parity tests for the +# compiled kernels live in unit_test.py. +# +# python test_index_math.py +# -------------------------------------------------------- + +import unittest + +import torch + +import reference as ref + + +# (H, W, window_h, window_w) -- covers square and non-square feature maps, and +# square and non-square windows. +SHAPES = [ + (8, 8, 4, 4), # nH == nW == 2 square grid, square window + (16, 8, 4, 4), # nH=4, nW=2 taller than wide + (8, 16, 4, 4), # nH=2, nW=4 wider than tall + (8, 32, 4, 8), # nH=2, nW=4 non-square window + (32, 8, 8, 4), # nH=4, nW=2 non-square window, taller than wide + (24, 24, 4, 8), # nH=6, nW=3 square grid, non-square window + (12, 8, 4, 4), # nH=3, nW=2 coprime counts, both > 1 + (12, 20, 4, 4), # nH=3, nW=5 coprime counts, nH < nW + (8, 8, 8, 8), # nH == nW == 1 single window + (12, 12, 4, 4), # nH == nW == 3 odd window count + (16, 32, 4, 8), # nH == nW == 4 square grid, non-square window + (4, 16, 4, 4), # nH=1, nW=4 one row of windows: the bug hides here + (16, 4, 4, 4), # nH=4, nW=1 one column of windows: it does not +] + +BATCH = 2 +CHANNELS = 3 + + +def _shifts(window_h, window_w): + """W-MSA (no shift), SW-MSA (half window), an asymmetric shift_h != shift_w + that only the per-axis signature on this branch can express, and the + negative form: models/swin_transformer.py calls the partition path with + -shift_size, so a suite that only ever passes positive shifts leaves the + sign the model actually uses untested.""" + return [(0, 0), (window_h // 2, window_w // 2), (1, 2), + (-(window_h // 2), -(window_w // 2))] + + +def _make_spatial(B, H, W, C): + return torch.arange(B * H * W * C, dtype=torch.float32).view(B, H, W, C) + + +def _make_windows(B, H, W, C, window_h, window_w): + n = B * (H // window_h) * (W // window_w) + return torch.arange(n * window_h * window_w * C, dtype=torch.float32).view( + n, window_h, window_w, C) + + +class TestKernelIndexMath(unittest.TestCase): + """Each kernel must equal the composition of PyTorch ops it replaces. + + The kernel `shift_h`/`shift_w` parameters are the negated torch.roll shift on + the partition path: callers pass -shift_size to the forward kernels and + +shift_size to the reverse ones. These identities pin that convention down. + """ + + def test_k1_equals_roll_then_partition(self): + for H, W, wh, ww in SHAPES: + for sh, sw in _shifts(wh, ww): + with self.subTest(H=H, W=W, wh=wh, ww=ww, shift=(sh, sw)): + x = _make_spatial(BATCH, H, W, CHANNELS) + got = ref.roll_and_window_partition_forward(x, (sh, sw), (wh, ww)) + want = ref.window_partition( + torch.roll(x, shifts=(sh, sw), dims=(1, 2)), (wh, ww)) + self.assertTrue(torch.equal(got, want)) + + def test_k2_equals_reverse_then_unroll(self): + for H, W, wh, ww in SHAPES: + for sh, sw in _shifts(wh, ww): + with self.subTest(H=H, W=W, wh=wh, ww=ww, shift=(sh, sw)): + g = _make_windows(BATCH, H, W, CHANNELS, wh, ww) + got = ref.roll_and_window_partition_backward(g, (sh, sw), (wh, ww), H, W) + want = torch.roll( + ref.window_reverse(g, (wh, ww), H, W), shifts=(-sh, -sw), dims=(1, 2)) + self.assertTrue(torch.equal(got, want)) + + def test_k3_equals_reverse_then_roll(self): + for H, W, wh, ww in SHAPES: + for sh, sw in _shifts(wh, ww): + with self.subTest(H=H, W=W, wh=wh, ww=ww, shift=(sh, sw)): + x = _make_windows(BATCH, H, W, CHANNELS, wh, ww) + got = ref.window_merge_and_roll_forward(x, (sh, sw), (wh, ww), H, W) + want = torch.roll( + ref.window_reverse(x, (wh, ww), H, W), shifts=(sh, sw), dims=(1, 2)) + self.assertTrue(torch.equal(got, want)) + + def test_k4_equals_unroll_then_partition(self): + for H, W, wh, ww in SHAPES: + for sh, sw in _shifts(wh, ww): + with self.subTest(H=H, W=W, wh=wh, ww=ww, shift=(sh, sw)): + g = _make_spatial(BATCH, H, W, CHANNELS) + got = ref.window_merge_and_roll_backward(g, (sh, sw), (wh, ww)) + want = ref.window_partition( + torch.roll(g, shifts=(-sh, -sw), dims=(1, 2)), (wh, ww)) + self.assertTrue(torch.equal(got, want)) + + def test_backward_kernels_invert_their_forward(self): + """K1/K2 and K3/K4 are permutations, so each pair must round-trip.""" + for H, W, wh, ww in SHAPES: + for sh, sw in _shifts(wh, ww): + with self.subTest(H=H, W=W, wh=wh, ww=ww, shift=(sh, sw)): + x = _make_spatial(BATCH, H, W, CHANNELS) + k1 = ref.roll_and_window_partition_forward(x, (sh, sw), (wh, ww)) + self.assertTrue(torch.equal( + ref.roll_and_window_partition_backward(k1, (sh, sw), (wh, ww), H, W), x)) + + w = _make_windows(BATCH, H, W, CHANNELS, wh, ww) + k3 = ref.window_merge_and_roll_forward(w, (sh, sw), (wh, ww), H, W) + self.assertTrue(torch.equal( + ref.window_merge_and_roll_backward(k3, (sh, sw), (wh, ww)), w)) + + +class TestUpstreamRowStrideBug(unittest.TestCase): + """window_merge_and_roll_forward used `* nH` where the row stride is `* nW`. + + Windows are laid out row-major as `b * nH * nW + wrow * nW + wcol`, so the + stride between consecutive window rows is nW. The two agree exactly when + nH == nW, which is why the bug never surfaced: every model in this repository + is trained on square images. + """ + + def _legacy_index(self, B, H, W, sh, sw, wh, ww): + return ref.window_merge_and_roll_forward_index( + B, H, W, sh, sw, wh, ww, legacy_row_stride=True) + + def test_legacy_is_correct_only_on_square_grids(self): + for H, W, wh, ww in SHAPES: + nH, nW = H // wh, W // ww + if nH != nW: + continue + for sh, sw in _shifts(wh, ww): + with self.subTest(H=H, W=W, wh=wh, ww=ww, shift=(sh, sw), nH=nH, nW=nW): + x = _make_windows(BATCH, H, W, CHANNELS, wh, ww) + want = ref.window_merge_and_roll_forward(x, (sh, sw), (wh, ww), H, W) + legacy = ref.window_merge_and_roll_forward( + x, (sh, sw), (wh, ww), H, W, legacy_row_stride=True) + self.assertTrue(torch.equal(legacy, want)) + + def test_legacy_reads_out_of_bounds_when_nH_greater_than_nW(self): + """nH > nW makes the miscomputed offset exceed the input, an illegal read.""" + checked = 0 + for H, W, wh, ww in SHAPES: + nH, nW = H // wh, W // ww + if nH <= nW: + continue + for sh, sw in _shifts(wh, ww): + with self.subTest(H=H, W=W, wh=wh, ww=ww, shift=(sh, sw), nH=nH, nW=nW): + numel = BATCH * nH * nW * wh * ww + legacy_max = int(self._legacy_index(BATCH, H, W, sh, sw, wh, ww).max()) + self.assertGreaterEqual(legacy_max, numel) + checked += 1 + self.assertGreater(checked, 0, "no nH > nW shape exercised") + + def test_legacy_returns_wrong_values_when_nH_less_than_nW(self): + """nH < nW keeps the offset in bounds, so the corruption is silent.""" + checked = 0 + for H, W, wh, ww in SHAPES: + nH, nW = H // wh, W // ww + if nH >= nW or nH == 1: + continue # nH == 1 hides the bug -- covered separately + for sh, sw in _shifts(wh, ww): + with self.subTest(H=H, W=W, wh=wh, ww=ww, shift=(sh, sw), nH=nH, nW=nW): + numel = BATCH * nH * nW * wh * ww + legacy_max = int(self._legacy_index(BATCH, H, W, sh, sw, wh, ww).max()) + self.assertLess(legacy_max, numel) + + x = _make_windows(BATCH, H, W, CHANNELS, wh, ww) + want = ref.window_merge_and_roll_forward(x, (sh, sw), (wh, ww), H, W) + legacy = ref.window_merge_and_roll_forward( + x, (sh, sw), (wh, ww), H, W, legacy_row_stride=True) + self.assertFalse(torch.equal(legacy, want)) + checked += 1 + self.assertGreater(checked, 0, "no nH < nW shape exercised") + + def test_legacy_is_invisible_when_nH_is_one(self): + """A single row of windows hides the bug entirely; a single column does not. + + The legacy term is `src_y / window_h * nH`, and `src_y / window_h` runs + over [0, nH). At nH == 1 it is always 0, so the wrong stride is never + multiplied by anything and the result is bit-identical to the correct + one -- even though nH != nW and the grid is as non-square as it gets. + The transpose has no such reprieve: nW == 1 with nH > 1 is the + out-of-bounds case, asserted above. + + This is why H != W is not on its own enough to reproduce the bug, and + why the shape list has to contain both orientations. + """ + checked = 0 + for H, W, wh, ww in SHAPES: + nH, nW = H // wh, W // ww + if nH != 1 or nW == 1: + continue + for sh, sw in _shifts(wh, ww): + with self.subTest(H=H, W=W, wh=wh, ww=ww, shift=(sh, sw), nH=nH, nW=nW): + x = _make_windows(BATCH, H, W, CHANNELS, wh, ww) + want = ref.window_merge_and_roll_forward(x, (sh, sw), (wh, ww), H, W) + legacy = ref.window_merge_and_roll_forward( + x, (sh, sw), (wh, ww), H, W, legacy_row_stride=True) + self.assertTrue(torch.equal(legacy, want)) + checked += 1 + self.assertGreater(checked, 0, "no nH == 1 < nW shape exercised") + + +class TestIntraWindowModuloIsEquivalent(unittest.TestCase): + """The upstream intra-window modulo omits `% H` / `% W`, which is a no-op here. + + `(y - s + H) % window_h` and `((y - s + H) % H) % window_h` differ by a + multiple of H, and H is a multiple of window_h whenever the launcher's + `nH = H / window_h` is exact -- which it must be for the kernel to be valid + at all. The explicit form is kept for readability, not for correctness. + """ + + def test_forms_agree_under_exact_divisibility(self): + for H, W, wh, ww in SHAPES: + for sh, sw in _shifts(wh, ww): + with self.subTest(H=H, W=W, wh=wh, ww=ww, shift=(sh, sw)): + explicit = ref.window_merge_and_roll_forward_index( + BATCH, H, W, sh, sw, wh, ww, legacy_intra_modulo=False) + legacy = ref.window_merge_and_roll_forward_index( + BATCH, H, W, sh, sw, wh, ww, legacy_intra_modulo=True) + self.assertTrue(torch.equal(explicit, legacy)) + + +if __name__ == '__main__': + unittest.main(verbosity=2) diff --git a/kernels/window_process/test_model_parity.py b/kernels/window_process/test_model_parity.py new file mode 100644 index 000000000..081ea3abb --- /dev/null +++ b/kernels/window_process/test_model_parity.py @@ -0,0 +1,192 @@ +# -------------------------------------------------------- +# Fused kernel for window process for SwinTransformer +# Copyright (c) 2022 Nvidia +# Licensed under The MIT License [see LICENSE for details] +# -------------------------------------------------------- +# End-to-end check: a SwinTransformer must produce identical output with and +# without the fused window kernels. The kernels replace torch.roll plus +# window_partition, which are exact data movements, so the two paths are not +# merely close -- they are bit-for-bit equal. +# +# python test_model_parity.py +# -------------------------------------------------------- + +import os +import sys +import unittest + +import torch + +sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), '..', '..'))) + +try: + import swin_window_process # noqa: F401 + EXTENSION_AVAILABLE = True +except ImportError: + EXTENSION_AVAILABLE = False + +try: + from models.swin_transformer import SwinTransformer + MODEL_AVAILABLE = True +except ImportError: # timm missing, for instance + SwinTransformer = None + MODEL_AVAILABLE = False + + +# How far a gradient may move under the fused path before the test calls it a +# regression, as a multiple of the run-to-run noise measured on the same machine. +# Only gradients that already fail to reproduce against a second eager run are +# held to this; every other one must still be bit-exact. The noise floor is a +# single sample, hence the headroom: on an MI300X it measured 1.9e-9 against a +# gradient scale of 0.254, so even 8x of it stays seven orders of magnitude below +# the signal. +NOISE_HEADROOM = 8 + +requires_everything = unittest.skipIf( + not (torch.cuda.is_available() and EXTENSION_AVAILABLE and MODEL_AVAILABLE), + 'requires a CUDA device, the compiled extension and an importable model') + + +def set_fused(model, enabled): + """Flip the flag on every block, so both runs share exactly the same weights.""" + switched = 0 + for module in model.modules(): + if hasattr(module, 'fused_window_process'): + module.fused_window_process = enabled + switched += 1 + return switched + + +@requires_everything +class TestModelParity(unittest.TestCase): + + def _model(self, img_size=56, window_size=7, depths=(2, 2), num_heads=(3, 6)): + torch.manual_seed(0) + model = SwinTransformer( + img_size=img_size, + patch_size=4, + in_chans=3, + num_classes=10, + embed_dim=48, + depths=list(depths), + num_heads=list(num_heads), + window_size=window_size, + drop_path_rate=0.0, + fused_window_process=False, + ) + return model.cuda().eval() + + def _assert_non_square_parity(self, img_size): + """img_size is documented as `int | tuple(int)`, and a non-square one + gives nH != nW inside the shifted blocks -- the configuration the row + stride fix is about. + + For img_size=(256, 128) with window_size=8 the shifted blocks run at + 64x32 and 32x16, so nH/nW are 8/4 and 4/2; the transpose (128, 256) + gives nW > nH instead. Before the fix the merge kernel mis-indexes on + both -- out of bounds when nH > nW, silently wrong when nH < nW. + """ + H, W = img_size + model = self._model(img_size=img_size, window_size=8, + depths=(2, 2, 2), num_heads=(3, 6, 12)) + + shifted_non_square = [ + m for m in model.modules() + if type(m).__name__ == 'SwinTransformerBlock' + and m.shift_size > 0 + and (m.input_resolution[0] // m.window_size) + != (m.input_resolution[1] // m.window_size) + ] + self.assertGreater(len(shifted_non_square), 0, + 'this configuration must exercise nH != nW') + + x = torch.randn(1, 3, H, W, device='cuda') + with torch.no_grad(): + eager = model(x) + set_fused(model, True) + fused = model(x) + set_fused(model, False) + + self.assertTrue(torch.equal(eager, fused)) + + def test_non_square_image_tall(self): + """H > W: the shifted blocks hit nH > nW, the out-of-bounds case.""" + self._assert_non_square_parity((256, 128)) + + def test_non_square_image_wide(self): + """W > H: the shifted blocks hit nH < nW, the silent-corruption case.""" + self._assert_non_square_parity((128, 256)) + + def test_forward_is_identical(self): + model = self._model() + x = torch.randn(2, 3, 56, 56, device='cuda') + + with torch.no_grad(): + eager = model(x) + blocks = set_fused(model, True) + fused = model(x) + set_fused(model, False) + + # set_fused counts every block carrying the flag, not just the shifted + # ones that route through the kernels. Zero would mean the model exposes + # no flag at all, and that the comparison below proved nothing. + self.assertGreater(blocks, 0, 'no block exposes fused_window_process') + self.assertTrue(torch.equal(eager, fused)) + + def test_gradients_are_identical(self): + """The fused path must not perturb a whole-model backward. + + Bit-exactness of the four kernels themselves is asserted in + unit_test.py, where it holds for every dtype and shape. At model level + the bar has to account for the rest of the network: the backward of a + GEMM is not always reproducible run to run, because the library is free + to pick a different reduction split each time. On an MI300X one of the + 63 gradients moves by ~2e-9 between two *identical* eager runs, so + asserting torch.equal against eager would fail without any fused kernel + being involved. + + So the eager path is run twice first, and each gradient is held to what + that measurement licenses: bit-exactness wherever eager reproduces + itself, and no further from eager than eager is from itself elsewhere. + On a platform where the backward is fully deterministic -- CUDA, in + every run of this test so far -- every gradient takes the first branch + and this is exactly the strict comparison it replaces. + """ + model = self._model() + x = torch.randn(2, 3, 56, 56, device='cuda') + target = torch.randn(2, 10, device='cuda') + + def grads(): + model.zero_grad(set_to_none=True) + torch.nn.functional.mse_loss(model(x), target).backward() + return [p.grad.clone() for p in model.parameters() if p.grad is not None] + + eager = grads() + eager_again = grads() # the platform's own run-to-run noise + set_fused(model, True) + fused = grads() + set_fused(model, False) + + self.assertEqual(len(eager), len(fused)) + reproducible = 0 + for i, (a, a2, b) in enumerate(zip(eager, eager_again, fused)): + with self.subTest(parameter=i): + if torch.equal(a, a2): + reproducible += 1 + self.assertTrue(torch.equal(a, b)) + else: + noise = (a - a2).abs().max().item() + delta = (a - b).abs().max().item() + self.assertLessEqual( + delta, NOISE_HEADROOM * noise, + f'gradient {i} moves {delta:.3e} with the fused path, ' + f'against {noise:.3e} between two eager runs') + + # A platform that reproduces nothing would make this test vacuous. + self.assertGreater(reproducible, len(eager) // 2) + + +if __name__ == '__main__': + if not (torch.cuda.is_available() and EXTENSION_AVAILABLE and MODEL_AVAILABLE): + print('Skipping: needs a CUDA device, the built extension and the model.\n') + unittest.main(verbosity=2) diff --git a/kernels/window_process/unit_test.py b/kernels/window_process/unit_test.py index 65dee5661..52880ef5c 100644 --- a/kernels/window_process/unit_test.py +++ b/kernels/window_process/unit_test.py @@ -3,248 +3,316 @@ # Copyright (c) 2022 Nvidia # Licensed under The MIT License [see LICENSE for details] # -------------------------------------------------------- +# Parity tests for the compiled kernels against the PyTorch ops they replace. +# Requires a CUDA device and the built extension; the index math itself is +# covered without either by test_index_math.py. +# +# python unit_test.py +# -------------------------------------------------------- -import torch -import swin_window_process -import random -import time import unittest +import torch -class WindowProcess(torch.autograd.Function): - @staticmethod - def forward(ctx, input, B, H, W, C, shift_size, window_size): - output = swin_window_process.roll_and_window_partition_forward(input, B, H, W, C, shift_size, window_size) - - ctx.B = B - ctx.H = H - ctx.W = W - ctx.C = C - ctx.shift_size = shift_size - ctx.window_size = window_size - return output - - @staticmethod - def backward(ctx, grad_in): - B = ctx.B - H = ctx.H - W = ctx.W - C = ctx.C - shift_size = ctx.shift_size - window_size = ctx.window_size - - grad_out = swin_window_process.roll_and_window_partition_backward(grad_in, B, H, W, C, shift_size, window_size) - return grad_out, None, None, None, None, None, None, None - - -class WindowProcessReverse(torch.autograd.Function): - @staticmethod - def forward(ctx, input, B, H, W, C, shift_size, window_size): - output = swin_window_process.window_merge_and_roll_forward(input, B, H, W, C, shift_size, window_size) - - ctx.B = B - ctx.H = H - ctx.W = W - ctx.C = C - ctx.shift_size = shift_size - ctx.window_size = window_size - - return output - - @staticmethod - def backward(ctx, grad_in): - B = ctx.B - H = ctx.H - W = ctx.W - C = ctx.C - shift_size = ctx.shift_size - window_size = ctx.window_size - - grad_out = swin_window_process.window_merge_and_roll_backward(grad_in, B, H, W, C, shift_size, window_size) - return grad_out, None, None, None, None, None, None, None - - -def window_partition(x, window_size): - """ - Args: - x: (B, H, W, C) - window_size (int): window size - Returns: - windows: (num_windows*B, window_size, window_size, C) - """ - B, H, W, C = x.shape - x = x.view(B, H // window_size, window_size, W // window_size, window_size, C) - windows = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, window_size, window_size, C) - return windows +import reference as ref + +try: + from window_process import WindowProcess, WindowProcessReverse + EXTENSION_AVAILABLE = True +except ImportError: # extension not built here + WindowProcess = WindowProcessReverse = None + EXTENSION_AVAILABLE = False + + +CUDA_AVAILABLE = torch.cuda.is_available() +requires_cuda = unittest.skipIf( + not (CUDA_AVAILABLE and EXTENSION_AVAILABLE), + 'requires a CUDA device and the compiled swin_window_process extension') + + +def available_dtypes(): + """float64 and float32 come from AT_DISPATCH_FLOATING_TYPES; the rest are added.""" + dtypes = [torch.float64, torch.float32, torch.float16] + if CUDA_AVAILABLE and torch.cuda.is_bf16_supported(): + dtypes.append(torch.bfloat16) + return dtypes + + +# (B, H, W, C, window_h, window_w) +SHAPES = [ + (2, 56, 56, 96, 7, 7), # nH == nW == 8 the ImageNet configuration + (2, 32, 16, 64, 8, 8), # nH=4, nW=2 out-of-bounds read before the fix + (2, 16, 32, 64, 8, 8), # nH=2, nW=4 silent corruption before the fix + (2, 8, 32, 64, 8, 8), # nH=1, nW=4 one row of windows: hides the bug + (2, 16, 64, 64, 4, 16), # nH=4, nW=4 non-square window + (2, 64, 16, 64, 16, 4), # nH=4, nW=4 non-square window, transposed + (2, 48, 48, 32, 6, 16), # nH=8, nW=3 square image, non-square window + (2, 24, 16, 64, 8, 8), # nH=3, nW=2 coprime counts, both > 1 + (2, 24, 40, 32, 8, 8), # nH=3, nW=5 coprime counts, nH < nW + (2, 24, 24, 32, 24, 24), # nH == nW == 1 a single window + # C selects the block width in best_block_dim(): < 384 -> 64 threads, + # < 1024 -> 128, otherwise 256. All three branches must be exercised, and C + # must be smaller than, equal to and larger than the block width so that the + # strided `for (i = threadIdx.x; i < C; i += blockDim.x)` loop is covered in + # all four of its forms: no full pass, exactly one, an exact number, and a + # ragged one where the last pass is partial. + (2, 16, 16, 1, 8, 8), # C < blockDim 64 threads, 63 idle + (2, 16, 16, 64, 8, 8), # C == blockDim 64 threads, one pass + (2, 16, 64, 100, 4, 16), # C % blockDim != 0 64 threads, 2 passes, 36 wide, + # and a non-square window with it + (2, 16, 16, 512, 8, 8), # 384 <= C < 1024 128 threads, 4 exact passes + (2, 16, 16, 1024, 8, 8), # C >= 1024 256 threads, 4 exact passes +] + + +def shifts_for(window_h, window_w): + """W-MSA (no shift), SW-MSA (half window, the isotropic shift Swin uses), and + an asymmetric shift_h != shift_w -- a regime only the per-axis signature on + this branch can express, so it needs its own coverage.""" + return [(0, 0), (window_h // 2, window_w // 2), (1, 2)] + + +def pyt_forward(x, shift, window): + """torch.roll followed by window_partition -- what WindowProcess fuses.""" + shift_h, shift_w = ref.to_pair(shift) + if shift_h or shift_w: + x = torch.roll(x, shifts=(-shift_h, -shift_w), dims=(1, 2)) + return ref.window_partition(x, window) + + +def reverse_pyt_forward(windows, shift, window, H, W): + """window_reverse followed by torch.roll -- what WindowProcessReverse fuses.""" + shift_h, shift_w = ref.to_pair(shift) + x = ref.window_reverse(windows, window, H, W) + if shift_h or shift_w: + x = torch.roll(x, shifts=(shift_h, shift_w), dims=(1, 2)) + return x + + +def leaf(tensor, requires_grad=True): + # .cuda() must come before requires_grad_: moving a tensor that already + # requires grad makes the result a non-leaf, and .grad never populates on it. + return tensor.clone().detach().cuda().requires_grad_(requires_grad) + + +def spatial(B, H, W, C, dtype): + return torch.randn((B, H, W, C), dtype=dtype) + + +def windowed(B, H, W, C, window_h, window_w, dtype): + n = B * (H // window_h) * (W // window_w) + return torch.randn((n, window_h, window_w, C), dtype=dtype) + + +@requires_cuda +class TestWindowProcess(unittest.TestCase): + """The kernels are exact permutations, so parity must be bit-for-bit. -def window_reverse(windows, window_size, H, W): + torch.equal is a stronger check than gradcheck here: no arithmetic is + performed on the values, so any deviation is an indexing error, not a + numerical one. Every dtype is therefore held to exact equality. """ - Args: - windows: (num_windows*B, window_size, window_size, C) - window_size (int): Window size - H (int): Height of image - W (int): Width of image - Returns: - x: (B, H, W, C) + + def _cases(self): + for dtype in available_dtypes(): + for B, H, W, C, window_h, window_w in SHAPES: + for shift in shifts_for(window_h, window_w): + yield dtype, B, H, W, C, (window_h, window_w), shift + + def test_partition_forward(self): + for dtype, B, H, W, C, window, shift in self._cases(): + with self.subTest(dtype=dtype, shape=(B, H, W, C), window=window, shift=shift): + x = spatial(B, H, W, C, dtype) + with torch.no_grad(): + expected = pyt_forward(leaf(x), shift, window) + fused = WindowProcess.apply( + leaf(x), B, H, W, C, (-shift[0], -shift[1]), window) + self.assertTrue(torch.equal(expected, fused)) + + def test_partition_backward(self): + for dtype, B, H, W, C, window, shift in self._cases(): + with self.subTest(dtype=dtype, shape=(B, H, W, C), window=window, shift=shift): + x = spatial(B, H, W, C, dtype) + grad = windowed(B, H, W, C, window[0], window[1], dtype).cuda() + + a, b = leaf(x), leaf(x) + pyt_forward(a, shift, window).backward(grad) + WindowProcess.apply( + b, B, H, W, C, (-shift[0], -shift[1]), window).backward(grad) + + self.assertIsNotNone(a.grad) + self.assertTrue(torch.equal(a.grad, b.grad)) + + def test_merge_forward(self): + for dtype, B, H, W, C, window, shift in self._cases(): + with self.subTest(dtype=dtype, shape=(B, H, W, C), window=window, shift=shift): + x = windowed(B, H, W, C, window[0], window[1], dtype) + with torch.no_grad(): + expected = reverse_pyt_forward(leaf(x), shift, window, H, W) + fused = WindowProcessReverse.apply(leaf(x), B, H, W, C, shift, window) + self.assertTrue(torch.equal(expected, fused)) + + def test_merge_backward(self): + for dtype, B, H, W, C, window, shift in self._cases(): + with self.subTest(dtype=dtype, shape=(B, H, W, C), window=window, shift=shift): + x = windowed(B, H, W, C, window[0], window[1], dtype) + grad = spatial(B, H, W, C, dtype).cuda() + + a, b = leaf(x), leaf(x) + reverse_pyt_forward(a, shift, window, H, W).backward(grad) + WindowProcessReverse.apply(b, B, H, W, C, shift, window).backward(grad) + + self.assertIsNotNone(a.grad) + self.assertTrue(torch.equal(a.grad, b.grad)) + + def test_round_trip(self): + """WindowProcessReverse must undo WindowProcess for the same shift.""" + for dtype, B, H, W, C, window, shift in self._cases(): + with self.subTest(dtype=dtype, shape=(B, H, W, C), window=window, shift=shift): + x = leaf(spatial(B, H, W, C, dtype), requires_grad=False) + windows = WindowProcess.apply( + x, B, H, W, C, (-shift[0], -shift[1]), window) + back = WindowProcessReverse.apply( + windows.contiguous(), B, H, W, C, shift, window) + self.assertTrue(torch.equal(x, back)) + + + def test_non_contiguous_gradient(self): + """An upstream op can hand the backward a non-contiguous gradient. + + The C++ side asserts contiguity, so window_process.py normalises it. + + The window is non-square on purpose: with 8x8 the transpose below is + shape-invariant, so this would pass even if the two window axes were + swapped somewhere along the path. + """ + B, H, W, C, window, shift = 2, 16, 16, 32, (4, 16), (2, 8) + n = B * (H // window[0]) * (W // window[1]) + x = spatial(B, H, W, C, torch.float32) + grad = torch.randn((n, window[1], window[0], C)).cuda().transpose(1, 2) + self.assertFalse(grad.is_contiguous()) + + a, b = leaf(x), leaf(x) + pyt_forward(a, shift, window).backward(grad) + WindowProcess.apply(b, B, H, W, C, (-shift[0], -shift[1]), window).backward(grad) + self.assertTrue(torch.equal(a.grad, b.grad)) + + +@requires_cuda +class TestStreamSemantics(unittest.TestCase): + """The kernels must run on the current stream, not on the default one. + + A race against the default stream is timing dependent and makes a poor + test. CUDA graph capture is the deterministic form of the same property: + capture runs on a non-default stream and rejects any kernel launched on the + legacy default stream outright, so this fails to capture unless the launch + uses at::cuda::getCurrentCUDAStream(). """ - B = int(windows.shape[0] / (H * W / window_size / window_size)) - x = windows.view(B, H // window_size, W // window_size, window_size, window_size, -1) - x = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(B, H, W, -1) - return x + def test_capturable_in_a_cuda_graph(self): + B, H, W, C, window, shift = 2, 16, 16, 32, (8, 8), (4, 4) + static_in = torch.randn((B, H, W, C), device='cuda') + + def call(): + return WindowProcess.apply( + static_in, B, H, W, C, (-shift[0], -shift[1]), window) + + # Warm up on a side stream, as the CUDA graph capture protocol requires. + side = torch.cuda.Stream() + side.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(side): + for _ in range(3): + call() + torch.cuda.current_stream().wait_stream(side) + + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + static_out = call() + graph.replay() + torch.cuda.synchronize() -def pyt_forward(x, shift_size, window_size): - # x in shape(B, H, W, C) - # cyclic shift - if shift_size > 0: - shifted_x = torch.roll(x, shifts=(-shift_size, -shift_size), dims=(1, 2)) - else: - shifted_x = x - # partition windows - x_windows = window_partition(shifted_x, window_size) - return x_windows - - -def reverse_pyt_forward(attn_windows, shift_size, window_size, H, W): - # x in shape(B*nH*nW, window_size, window_size, C) - shifted_x = window_reverse(attn_windows, window_size, H, W) - if shift_size > 0: - x = torch.roll(shifted_x, shifts=(shift_size, shift_size), dims=(1, 2)) - else: - x = shifted_x - return x + self.assertTrue(torch.equal(static_out, pyt_forward(static_in, shift, window))) + def test_matches_when_run_on_a_side_stream(self): + B, H, W, C, window, shift = 2, 16, 16, 32, (8, 8), (4, 4) + x = torch.randn((B, H, W, C), device='cuda') + expected = pyt_forward(x, shift, window) -def copy_one_tensor(input, requires_grad=True): - input1 = input.clone().detach().requires_grad_(requires_grad).cuda() - return input1 + side = torch.cuda.Stream() + side.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(side): + fused = WindowProcess.apply(x, B, H, W, C, (-shift[0], -shift[1]), window) + torch.cuda.current_stream().wait_stream(side) -class Test_WindowProcess(unittest.TestCase): - def setUp(self): - self.B = 192 - self.H = 56 - self.W = 56 - self.C = 96 - self.shift_size = 2 - self.window_size = 7 - self.nH = self.H // self.window_size - self.nW = self.W // self.window_size - - def test_roll_and_window_partition_forward(self, dtype=torch.float32): - input = torch.randn((self.B, self.H, self.W, self.C), dtype=dtype, requires_grad=True).cuda() - - input1 = copy_one_tensor(input, True) - input2 = copy_one_tensor(input, True) + self.assertTrue(torch.equal(expected, fused)) - with torch.no_grad(): - # ori - expected = pyt_forward(input1, self.shift_size, self.window_size) - # fused kernel - fused_output = WindowProcess.apply(input2, self.B, self.H, self.W, self.C, -self.shift_size, self.window_size) - - self.assertTrue(torch.equal(expected, fused_output)) - #self.assertTrue(torch.allclose(expected, fused_output, rtol=1e-05, atol=1e-08)) - - def test_roll_and_window_partition_backward(self, dtype=torch.float32): - input = torch.randn((self.B, self.H, self.W, self.C), dtype=dtype, requires_grad=True).cuda() - d_loss_tensor = torch.randn((self.B*self.nW*self.nH, self.window_size, self.window_size, self.C), dtype=dtype).cuda() - - input1 = copy_one_tensor(input, True) - input2 = copy_one_tensor(input, True) - - # ori - expected = pyt_forward(input1, self.shift_size, self.window_size) - expected.backward(d_loss_tensor) - # fused kernel - fused_output = WindowProcess.apply(input2, self.B, self.H, self.W, self.C, -self.shift_size, self.window_size) - fused_output.backward(d_loss_tensor) - - self.assertTrue(torch.equal(expected, fused_output)) - #self.assertTrue(torch.allclose(expected, fused_output, rtol=1e-05, atol=1e-08)) - - def test_window_merge_and_roll_forward(self, dtype=torch.float32): - input = torch.randn((self.B*self.nH*self.nW, self.window_size, self.window_size, self.C), dtype=dtype, requires_grad=True).cuda() - - input1 = copy_one_tensor(input, True) - input2 = copy_one_tensor(input, True) - with torch.no_grad(): - # ori - expected = reverse_pyt_forward(input1, self.shift_size, self.window_size, self.H, self.W) - # fused kernel - fused_output = WindowProcessReverse.apply(input2, self.B, self.H, self.W, self.C, self.shift_size, self.window_size) - - self.assertTrue(torch.equal(expected, fused_output)) - #self.assertTrue(torch.allclose(expected, fused_output, rtol=1e-05, atol=1e-08)) - - - def test_window_merge_and_roll_backward(self, dtype=torch.float32): - input = torch.randn((self.B*self.nH*self.nW, self.window_size, self.window_size, self.C), dtype=dtype, requires_grad=True).cuda() - d_loss_tensor = torch.randn((self.B, self.H, self.W, self.C), dtype=dtype, requires_grad=True).cuda() - - input1 = copy_one_tensor(input, True) - input2 = copy_one_tensor(input, True) - - # ori - expected = reverse_pyt_forward(input1, self.shift_size, self.window_size, self.H, self.W) - expected.backward(d_loss_tensor) - # fused kernel - fused_output = WindowProcessReverse.apply(input2, self.B, self.H, self.W, self.C, self.shift_size, self.window_size) - fused_output.backward(d_loss_tensor) - - self.assertTrue(torch.equal(expected, fused_output)) - #self.assertTrue(torch.allclose(expected, fused_output, rtol=1e-05, atol=1e-08)) - - def test_forward_backward_speed(self, dtype=torch.float32, times=1000): - input = torch.randn((self.B*self.nH*self.nW, self.window_size, self.window_size, self.C), dtype=dtype, requires_grad=True).cuda() - d_loss_tensor = torch.randn((self.B, self.H, self.W, self.C), dtype=dtype, requires_grad=True).cuda() - - input1 = copy_one_tensor(input, True) - input2 = copy_one_tensor(input, True) - - # SwinTransformer official - def run_pyt(t=1000): - for _ in range(t): - expected = reverse_pyt_forward(input1, self.shift_size, self.window_size, self.H, self.W) - expected.backward(d_loss_tensor) - - # my op - def run_fusedop(t=1000): - for _ in range(t): - fused_output = WindowProcessReverse.apply(input2, self.B, self.H, self.W, self.C, self.shift_size, self.window_size) - fused_output.backward(d_loss_tensor) - - torch.cuda.synchronize() - t1 = time.time() - run_pyt(t=times) - torch.cuda.synchronize() - t2 = time.time() - run_fusedop(t=times) - torch.cuda.synchronize() - t3 = time.time() - self.assertTrue((t3 - t2) < (t2 - t1)) +@requires_cuda +class TestScalarWindowStillWorks(unittest.TestCase): + """An int shift/window must behave exactly as the (h, w) pair it expands to. - print('Run {} times'.format(times)) - print('Original time cost: {}'.format(t2 - t1)) - print('Fused op time cost: {}'.format(t3 - t2)) - - def test_roll_and_window_partition_forward_fp16(self, dtype=torch.float16): - self.test_roll_and_window_partition_forward(dtype=dtype) + This is the compatibility path used by models/swin_transformer.py, which + passes ints and must keep working unchanged. + """ + + def test_int_and_pair_agree(self): + B, H, W, C, window, shift = 2, 56, 56, 96, 7, 3 + x = spatial(B, H, W, C, torch.float32) + with torch.no_grad(): + as_int = WindowProcess.apply(leaf(x), B, H, W, C, -shift, window) + as_pair = WindowProcess.apply( + leaf(x), B, H, W, C, (-shift, -shift), (window, window)) + self.assertTrue(torch.equal(as_int, as_pair)) - def test_roll_and_window_partition_backward_fp16(self, dtype=torch.float16): - self.test_roll_and_window_partition_backward(dtype=dtype) - def test_window_merge_and_roll_forward_fp16(self, dtype=torch.float16): - self.test_window_merge_and_roll_forward(dtype=dtype) - - def test_window_merge_and_roll_backward_fp16(self, dtype=torch.float16): - self.test_window_merge_and_roll_backward(dtype=dtype) +@requires_cuda +class TestPreconditions(unittest.TestCase): + """Invalid arguments must raise instead of reading the wrong memory.""" - def test_forward_backward_speed_fp16(self, dtype=torch.float16, times=1000): - self.test_forward_backward_speed(dtype=dtype, times=times) + def setUp(self): + self.B, self.H, self.W, self.C = 2, 16, 16, 32 + self.x = leaf(spatial(self.B, self.H, self.W, self.C, torch.float32), False) + + def _apply(self, **overrides): + kwargs = dict(B=self.B, H=self.H, W=self.W, C=self.C, shift=(0, 0), window=(8, 8)) + kwargs.update(overrides) + return WindowProcess.apply( + self.x, kwargs['B'], kwargs['H'], kwargs['W'], kwargs['C'], + kwargs['shift'], kwargs['window']) + + def test_window_must_divide_height(self): + with self.assertRaisesRegex(RuntimeError, 'divisible by window_h'): + self._apply(window=(5, 8)) + + def test_window_must_divide_width(self): + with self.assertRaisesRegex(RuntimeError, 'divisible by window_w'): + self._apply(window=(8, 5)) + + def test_shift_must_be_smaller_than_window(self): + with self.assertRaisesRegex(RuntimeError, 'smaller than the window size'): + self._apply(shift=(8, 0)) + + def test_shape_must_match_tensor(self): + with self.assertRaisesRegex(RuntimeError, 'expected'): + self._apply(C=self.C * 2) + + def test_window_count_above_the_grid_limit_is_rejected(self): + """B * nH * nW is grid.z, which the CUDA driver caps at 65535.""" + B, H, W, C = 70000, 8, 8, 1 + big = leaf(spatial(B, H, W, C, torch.float32), False) + with self.assertRaisesRegex(RuntimeError, 'grid.z limit'): + WindowProcess.apply(big, B, H, W, C, (0, 0), (8, 8)) + + def test_non_contiguous_input_is_rejected(self): + transposed = self.x.transpose(1, 2) + with self.assertRaises(RuntimeError): + WindowProcess.apply( + transposed, self.B, self.W, self.H, self.C, (0, 0), (8, 8)) if __name__ == '__main__': - print('Pass only two tensors are exactly the same (using torch.equal).\n') + if not (CUDA_AVAILABLE and EXTENSION_AVAILABLE): + print('No CUDA device or extension not built: every test here will be skipped.') + print('Run test_index_math.py to check the index math without a GPU.\n') torch.manual_seed(0) unittest.main(verbosity=2) diff --git a/kernels/window_process/window_process.py b/kernels/window_process/window_process.py index ee43e9e97..f64c76c68 100644 --- a/kernels/window_process/window_process.py +++ b/kernels/window_process/window_process.py @@ -8,43 +8,89 @@ import swin_window_process +def to_pair(value): + """Accept an int (isotropic) or a (h, w) iterable and return a (h, w) tuple. + + Keeps the public signature of WindowProcess unchanged: `window_size=7` + behaves exactly as before, `window_size=(4, 8)` selects a rectangular window. + """ + if isinstance(value, int): + return value, value + h, w = value + return int(h), int(w) + + class WindowProcess(torch.autograd.Function): + """Fused torch.roll + window_partition. + + Args: + input: (B, H, W, C) contiguous CUDA tensor. float64, float32, + float16 and bfloat16 all dispatch. + B, H, W, C: shape of `input`. + shift_size: int or (shift_h, shift_w). This is the *negated* torch.roll + shift, matching the existing call sites in models/swin_transformer.py. + window_size: int or (window_h, window_w). Must divide H and W per axis. + + Returns: + (B * nH * nW, window_h, window_w, C) + """ + @staticmethod def forward(ctx, input, B, H, W, C, shift_size, window_size): - output = swin_window_process.roll_and_window_partition_forward(input, B, H, W, C, shift_size, window_size) + shift_h, shift_w = to_pair(shift_size) + window_h, window_w = to_pair(window_size) + output = swin_window_process.roll_and_window_partition_forward( + input, B, H, W, C, shift_h, shift_w, window_h, window_w) ctx.B = B ctx.H = H - ctx.W = W - ctx.C = C - ctx.shift_size = shift_size - ctx.window_size = window_size + ctx.W = W + ctx.C = C + ctx.shift_size = (shift_h, shift_w) + ctx.window_size = (window_h, window_w) return output @staticmethod def backward(ctx, grad_in): B = ctx.B H = ctx.H - W = ctx.W - C = ctx.C - shift_size = ctx.shift_size - window_size = ctx.window_size + W = ctx.W + C = ctx.C + shift_h, shift_w = ctx.shift_size + window_h, window_w = ctx.window_size - grad_out = swin_window_process.roll_and_window_partition_backward(grad_in, B, H, W, C, shift_size, window_size) - return grad_out, None, None, None, None, None, None, None + grad_out = swin_window_process.roll_and_window_partition_backward( + grad_in.contiguous(), B, H, W, C, shift_h, shift_w, window_h, window_w) + return grad_out, None, None, None, None, None, None class WindowProcessReverse(torch.autograd.Function): + """Fused window merge + torch.roll, the inverse of WindowProcess. + + Args: + input: (B * nH * nW, window_h, window_w, C) contiguous CUDA tensor. + float64, float32, float16 and bfloat16 all dispatch. + B, H, W, C: shape of the *output* feature map. + shift_size: int or (shift_h, shift_w), the torch.roll shift. + window_size: int or (window_h, window_w). Must divide H and W per axis. + + Returns: + (B, H, W, C) + """ + @staticmethod def forward(ctx, input, B, H, W, C, shift_size, window_size): - output = swin_window_process.window_merge_and_roll_forward(input, B, H, W, C, shift_size, window_size) + shift_h, shift_w = to_pair(shift_size) + window_h, window_w = to_pair(window_size) + output = swin_window_process.window_merge_and_roll_forward( + input, B, H, W, C, shift_h, shift_w, window_h, window_w) ctx.B = B ctx.H = H - ctx.W = W - ctx.C = C - ctx.shift_size = shift_size - ctx.window_size = window_size + ctx.W = W + ctx.C = C + ctx.shift_size = (shift_h, shift_w) + ctx.window_size = (window_h, window_w) return output @@ -52,12 +98,11 @@ def forward(ctx, input, B, H, W, C, shift_size, window_size): def backward(ctx, grad_in): B = ctx.B H = ctx.H - W = ctx.W - C = ctx.C - shift_size = ctx.shift_size - window_size = ctx.window_size - - #grad_out = ctx.saved_tensors[0] - #grad_out = torch.zeros((B, H, W, C), dtype=dtype).cuda() - grad_out = swin_window_process.window_merge_and_roll_backward(grad_in, B, H, W, C, shift_size, window_size) - return grad_out, None, None, None, None, None, None, None + W = ctx.W + C = ctx.C + shift_h, shift_w = ctx.shift_size + window_h, window_w = ctx.window_size + + grad_out = swin_window_process.window_merge_and_roll_backward( + grad_in.contiguous(), B, H, W, C, shift_h, shift_w, window_h, window_w) + return grad_out, None, None, None, None, None, None