[onnxruntime] WebGPU MatMulNBits 최적화: Subgroup Shuffle을 활용한 성능 향상
PR 링크: microsoft/onnxruntime#31703 상태: Merged | 변경: +66 / -21
들어가며
대규모 언어 모델(LLM)의 추론 성능은 행렬 곱셈 연산의 효율성에 크게 좌우됩니다. 특히 양자화된 행렬 곱셈(MatMulNBits)은 메모리 대역폭과 연산량을 줄여 성능 향상에 기여하지만, GPU 아키텍처의 특성을 제대로 활용하지 못하면 오히려 병목 현상이 발생할 수 있습니다. Microsoft의 ONNX Runtime 레포지토리에서 올라온 이 PR은 WebGPU 백엔드에서 MatMulNBits 연산의 'Wide-Tile' 커널에 Subgroup Shuffle을 도입하여 A 피연산자의 중복된 워크그룹 메모리 읽기를 제거하고, 다양한 하드웨어 환경에서의 호환성과 성능을 개선하는 것을 목표로 합니다.
이 글에서는 해당 PR의 코드 변경 사항을 상세히 분석하고, Subgroup Shuffle 도입이 왜 성능 향상으로 이어지는지, 그리고 이 최적화가 가지는 일반적인 교훈은 무엇인지 살펴보겠습니다.
코드 분석
1. onnxruntime/contrib_ops/webgpu/quantization/matmul_nbits.cc
이 파일에서는 MatMulNBitsWideTileProgram의 생성자 호출 부분이 수정되었습니다. 주요 변경점은 다음과 같습니다.
-
Workgroup Size 및 Tile M 계산 변경: 기존에는
tile_m_이workgroup_size / 8로 고정되어 있었으나, 이제는workgroup_size가 3D 워크그룹까지 고려하도록WorkgroupSizeX() * WorkgroupSizeY() * WorkgroupSizeZ()로 확장되었습니다. 또한,tile_m_계산 시 데이터 타입(f32는 16, f16은 32)에 따라 값이 달라지도록 수정되었습니다. 이는 f16이 f32보다 두 배 많은 요소를 레지스터에 패킹할 수 있기 때문입니다.Before:
- const uint32_t workgroup_size = WorkgroupSizeX() * WorkgroupSizeY(); - ORT_ENFORCE(tile_m_ == workgroup_size / 8, "tile_m must be workgroup_size / 8."); + const uint32_t workgroup_size = WorkgroupSizeX() * WorkgroupSizeY() * WorkgroupSizeZ(); + ORT_ENFORCE(tile_m_ % (workgroup_size / 8) == 0, "tile_m must be a multiple of workgroup_size / 8.");After:
+ const bool is_f16 = a->DataType() == DataTypeImpl::GetType<MLFloat16>(); + const uint32_t tile_m = is_f16 ? 2 * workgroup_size / 8 : workgroup_size / 8; -
Subgroup 지원 여부 및 최적 전략 선택:
subgroup_min_size파라미터가MatMulNBitsWideTileProgram생성자에 추가되었습니다. 이 값은 어댑터가 보장하는 최소 Subgroup 크기이며, Subgroup 지원 여부와 NVIDIA GPU인지 여부에 따라 결정됩니다. NVIDIA GPU의 경우, 최적화된 공유 메모리 브로드캐스트가 더 효율적이라는 실험 결과에 따라 Subgroup Shuffle 최적화를 제외합니다.Before:
- MatMulNBitsWideTileProgram program{has_zero_points, has_bias, has_weight_idx, has_weight_idx_indirect, tile_m, tile_n, static_cast<uint32_t>(nbits)}; + bool is_nvidia = context.AdapterInfo().vendor == std::string_view{"nvidia"}; + const uint32_t subgroup_min_size = (context.HasFeature(wgpu::FeatureName::Subgroups) && !is_nvidia) + ? context.AdapterInfo().subgroupMinSize + : 0u; + MatMulNBitsWideTileProgram program{has_zero_points, has_bias, has_weight_idx, has_weight_idx_indirect, tile_m, tile_n, static_cast<uint32_t>(nbits), subgroup_min_size};After:
- program.CacheHint(nbits, has_zero_points, has_bias, has_weight_idx, has_weight_idx_indirect); + program.CacheHint(nbits, has_zero_points, has_bias, has_weight_idx, has_weight_idx_indirect, subgroup_min_size);
2. onnxruntime/contrib_ops/webgpu/quantization/matmul_nbits.h
MatMulNBitsWideTileProgram 클래스에 subgroup_min_size_ 멤버 변수가 추가되었습니다. 이는 WGSL 템플릿에서 Subgroup 관련 로직을 제어하는 데 사용됩니다.
Before:
- class MatMulNBitsWideTileProgram final : public Program<MatMulNBitsWideTileProgram> {
- public:
- MatMulNBitsWideTileProgram(bool has_zero_points, bool has_bias, bool has_weight_idx, bool has_weight_idx_indirect, uint32_t tile_m, uint32_t tile_n, uint32_t nbits)
- : Program{"MatMulNBitsWideTile"}, has_zero_points_{has_zero_points}, has_bias_{has_bias}, has_weight_idx_{has_weight_idx}, has_weight_idx_indirect_{has_weight_idx_indirect}, tile_m_(tile_m), tile_n_(tile_n), nbits_(nbits) {}
+ class MatMulNBitsWideTileProgram final : public Program<MatMulNBitsWideTileProgram> {
+ public:
+ MatMulNBitsWideTileProgram(bool has_zero_points, bool has_bias, bool has_weight_idx, bool has_weight_idx_indirect, uint32_t tile_m, uint32_t tile_n, uint32_t nbits, uint32_t subgroup_min_size)
+ : Program{"MatMulNBitsWideTile"}, has_zero_points_{has_zero_points}, has_bias_{has_bias}, has_weight_idx_{has_weight_idx}, has_weight_idx_indirect_{has_weight_idx_indirect}, tile_m_(tile_m), tile_n_(tile_n), nbits_(nbits), subgroup_min_size_(subgroup_min_size) {}
Status GenerateShaderCode(ShaderHelper& sh) const override;
WEBGPU_PROGRAM_DEFINE_UNIFORM_VARIABLES({"Batch", ProgramUniformVariableDataType::Uint32},{"num_N_tile", ProgramUniformVariableDataType::Uint32},{"num_M_tile", ProgramUniformVariableDataType::Uint32},{"n_blocks_per_col", ProgramUniformVariableDataType::Uint32},{"zero_blocks_per_col", ProgramUniformVariableDataType::Uint32});
@@ -37,5 +37,6 @@ class MatMulNBitsWideTileProgram final : public Program<MatMulNBitsWideTileProgr
uint32_t tile_m_;
uint32_t tile_n_;
uint32_t nbits_;
+ uint32_t subgroup_min_size_;
};
class MatMulNBitsProgram final : public Program<MatMulNBitsProgram> {
3. onnxruntime/contrib_ops/webgpu/quantization/matmul_nbits_wide_tile.wgsl.template
이 파일은 실제 WebGPU 셰이더 코드를 정의하는 템플릿입니다. Subgroup Shuffle 로직이 핵심적인 변경 사항입니다.
-
A 피연산자 로딩 및 저장 방식 변경: 기존에는 각 워커 스레드가 A 피연산자의 타일을 워크그룹 메모리에 직접 로드했습니다. 하지만 이 방식은 동일한 A 타일 데이터를 여러 스레드가 중복해서 읽을 수 있다는 비효율성을 가집니다. 변경된 코드에서는 각 스레드가 필요한 A 데이터의 일부(vec)만 로드하고,
subgroupShuffle을 사용하여 다른 스레드들이 필요한 데이터를 공유하도록 합니다. A 데이터 타일은[KAVecSizeForBlock32 / 2u][kTileM][2u]형태로 저장되어, 두 개의 vec을 한 번의 연속적인 메모리 접근으로 읽을 수 있게 최적화되었습니다.Before (로딩 로직 일부):
- // Load `a` elements into workgroup memory, TileM x KAVecSizeForBlock32 (block32) - let a_row_idx = local_idx / KAVecSizeForBlock32; - let a_col_idx = local_idx % KAVecSizeForBlock32; - a_data_tile[a_row_idx][a_col_idx] = load_a(batch, - row + a_row_idx, - block_idx * KAVecSizeForBlock32 + a_col_idx); + // Load `a` elements into workgroup memory, kTileM x KAVecSizeForBlock32 (block32), + // stored as pairs of vecs along K. One pass covers workgroup_size_x / KAVecSizeForBlock32 + // rows, so loop when kTileM is wider than that. + for (var a_row_base = 0u; a_row_base < kTileM; a_row_base += workgroup_size_x / KAVecSizeForBlock32) { + let a_row_idx = a_row_base + local_idx / KAVecSizeForBlock32; + let a_col_idx = local_idx % KAVecSizeForBlock32; + a_data_tile[a_col_idx / 2u][a_row_idx][a_col_idx % 2u] = load_a(batch, + row + a_row_idx, + block_idx * KAVecSizeForBlock32 + a_col_idx); + } -
Subgroup Shuffle 로직 도입: 셰이더의 메인 루프(
$MAIN) 내에서subgroup_min_size값에 따라 세 가지 전략으로 나뉩니다:subgroup_min_size >= tile_m: 전체kTileM행을 한 번의subgroupShuffle로 처리합니다. 각 스레드는 자신의a_data_tile에서 필요한 행을 로드하고,subgroupShuffle을 통해 다른 스레드들이 접근할 수 있게 합니다.subgroup_min_size > 0 && tile_m % subgroup_min_size == 0:kTileM이subgroup_min_size의 배수일 경우,kTileM을subgroup_min_size크기의 밴드(chunk)로 나누어subgroupShuffle을 반복 적용합니다.else(Subgroups 미지원 또는 기타 경우):subgroupShuffle을 사용하지 않고, 각 스레드가 워크그룹 메모리에서 직접 데이터를 읽어 처리하는 기존 방식(direct reduction)을 사용합니다.
Subgroup Shuffle 적용 코드:
+#if subgroup_min_size >= tile_m
- let capped_sg_id = min(sg_id, kTileM - 1u); +#elif subgroup_min_size > 0 && tile_m % subgroup_min_size == 0
- let capped_sg_id = min(sg_id, subgroup_min_size - 1u); +#endif
- var results : array<output_element_t, kTileM>; for (var block_idx = 0u; block_idx < uniforms.n_blocks_per_col; block_idx++) {
- // Load
aelements into workgroup memory, TileM x KAVecSizeForBlock32 (block32) - let a_row_idx = local_idx / KAVecSizeForBlock32;
- let a_col_idx = local_idx % KAVecSizeForBlock32;
- a_data_tile[a_row_idx][a_col_idx] = load_a(batch,
-
row + a_row_idx, -
block_idx * KAVecSizeForBlock32 + a_col_idx);
-
// Load
aelements into workgroup memory, kTileM x KAVecSizeForBlock32 (block32), -
// stored as pairs of vecs along K. One pass covers workgroup_size_x / KAVecSizeForBlock32
-
// rows, so loop when kTileM is wider than that.
-
for (var a_row_base = 0u; a_row_base < kTileM; a_row_base += workgroup_size_x / KAVecSizeForBlock32) {
-
let a_row_idx = a_row_base + local_idx / KAVecSizeForBlock32; -
let a_col_idx = local_idx % KAVecSizeForBlock32; -
a_data_tile[a_col_idx / 2u][a_row_idx][a_col_idx % 2u] = load_a(batch, -
row + a_row_idx, -
block_idx * KAVecSizeForBlock32 + a_col_idx); -
} workgroupBarrier();
let b_row = col + local_idx; @@ -207,15 +221,33 @@ $MAIN { let zero_point = load_zero(b_row, block_idx, uniforms.N, uniforms.zero_blocks_per_col); let b_data = load_b(b_row, block_idx);
- for (var b_idx = 0u; b_idx < 4u; b_idx++) {
- for (var b_idx = 0u; b_idx < KAVecSizeForBlock32 / 2u; b_idx++) { let b_dequantized = dequantize(b_data[b_idx], zero_point, scale); +#if subgroup_min_size >= tile_m
-
// Adapter guarantees a subgroup wide enough to shuffle the full kTileM band at once. -
let a = a_data_tile[b_idx][capped_sg_id]; for (var m_idx = 0u; m_idx < kTileM; m_idx++) {
-
let a_data0 = a_data_tile[m_idx][b_idx * 2u]; -
let a_data1 = a_data_tile[m_idx][b_idx * 2u + 1u]; -
results[m_idx] += dot(a_data0, b_dequantized[0]) + -
dot(a_data1, b_dequantized[1]);
-
results[m_idx] += dot(subgroupShuffle(a[0], m_idx), b_dequantized[0]) + -
dot(subgroupShuffle(a[1], m_idx), b_dequantized[1]); -
}
+#elif subgroup_min_size > 0 && tile_m % subgroup_min_size == 0
-
// kTileM is wider than a single subgroup can shuffle, so it's split into -
// subgroup_min_size-wide bands, each shuffled within its own subgroup. -
for (var chunk = 0u; chunk < kTileM / subgroup_min_size; chunk++) { -
let m_offset = chunk * subgroup_min_size; -
let a = a_data_tile[b_idx][capped_sg_id + m_offset]; -
for (var m_idx = 0u; m_idx < subgroup_min_size; m_idx++) { -
results[m_idx + m_offset] += dot(subgroupShuffle(a[0], m_idx), b_dequantized[0]) + -
dot(subgroupShuffle(a[1], m_idx), b_dequantized[1]); -
} -
}
+#else
-
// No guaranteed subgroup support wide enough for shuffles: read each row directly. -
for (var m_idx = 0u; m_idx < kTileM; m_idx++) { -
let a = a_data_tile[b_idx][m_idx]; -
results[m_idx] += dot(a[0], b_dequantized[0]) + dot(a[1], b_dequantized[1]); -
}
+#endif } workgroupBarrier(); }
* **Out-of-Bounds Access 방지:**
리뷰 과정에서 `sg_id`가 `kTileM`을 초과할 경우 `a_data_tile` 접근 시 Out-of-Bounds(OOB)가 발생할 수 있다는 지적이 있었습니다. 이를 방지하기 위해 `capped_sg_id = min(sg_id, kTileM - 1u)`와 같이 `sg_id`를 클램핑하는 로직이 추가되었습니다. 이는 실제 사용되지 않는 레인에서 발생하는 OOB 로드를 방지하여 안정성을 높입니다.
**수정된 코드:**
```diff
+#if subgroup_min_size >= tile_m
+ let capped_sg_id = min(sg_id, kTileM - 1u);
+#elif subgroup_min_size > 0 && tile_m % subgroup_min_size == 0
+ let capped_sg_id = min(sg_id, subgroup_min_size - 1u);
+#endif
```
## 왜 이게 좋은가?
이 PR의 핵심적인 개선은 **A 피연산자의 워크그룹 메모리 읽기 중복 제거**입니다. 기존 방식에서는 각 워커 스레드가 A 행렬의 타일 데이터를 워크그룹 메모리에 로드해야 했습니다. 이는 동일한 데이터를 여러 스레드가 반복적으로 읽는 비효율을 초래했습니다. Subgroup Shuffle을 사용하면, 각 스레드는 자신이 필요한 A 데이터의 일부만 로드하고, Subgroup 내에서 `subgroupShuffle` 함수를 통해 다른 스레드들이 필요로 하는 데이터를 효율적으로 공유할 수 있습니다. 결과적으로 워크그룹 메모리 접근 횟수가 줄어들어 메모리 대역폭 병목 현상이 완화되고, 연산 속도가 향상됩니다.
PR 설명에 포함된 성능 데이터는 이러한 개선 효과를 명확하게 보여줍니다.
| max_length = 8K | Prefill Length | Default Prefill TPS | Opt TileM-16 Prefill TPS | Opt TileM-32 Prefill TPS | Improvement |
| :----------- | -------------: | ------------------: | -----------------------: | -----------------------: | ----------: |
| Panther Lake | 128 | 221.01 | 242.49 | 243.02 | 110% |
| Panther Lake | 1024 | 594.20 | 644.74 | 696.70 | 117% |
| Panther Lake | 4096 | 596.47 | 667.17 | 727.84 | 122% |
| Alder Lake | 128 | 46.76 | 86.13 | 113.09 | 242% |
| Alder Lake | 1024 | 44.72 | 94.43 | 129.06 | 289% |
특히 Alder Lake 아키텍처에서 최대 289%까지의 TPS(Transactions Per Second) 향상을 기록한 것은 이 최적화의 강력한 효과를 입증합니다. 이는 Subgroup Shuffle이 GPU 하드웨어의 병렬 처리 능력을 더욱 효과적으로 활용할 수 있게 해주기 때문입니다.
또한, 이 PR은 다양한 하드웨어 환경에서의 호환성을 높였습니다. Subgroup 지원 여부, Subgroup의 최소 크기, 그리고 NVIDIA GPU와 같이 특정 아키텍처의 특성을 고려하여 최적의 전략(Subgroup Shuffle 또는 Direct Reduction)을 동적으로 선택하도록 구현되었습니다. 이는 특정 하드웨어에만 국한되지 않고 더 넓은 범위의 WebGPU 지원 장치에서 성능 이점을 얻을 수 있도록 합니다.
**일반적인 교훈:**
1. **데이터 재사용 및 통신 최적화:** GPU 커널에서 데이터 재사용은 성능 향상의 핵심입니다. 워크그룹 내에서 데이터를 한 번만 로드하고 `Subgroup Shuffle`과 같은 하드웨어 기능을 활용하여 효율적으로 공유하는 것은 메모리 대역폭을 절약하는 강력한 방법입니다.
2. **하드웨어 특성 고려:** GPU 아키텍처(예: Subgroup 지원 여부, NVIDIA의 공유 메모리 최적화)는 성능에 큰 영향을 미칩니다. 가능한 경우, 하드웨어의 특성을 파악하고 이에 맞는 최적화 전략을 적용하는 것이 중요합니다.
3. **동적 분기 및 적응성:** 다양한 하드웨어 및 환경에서 최적의 성능을 내기 위해, 런타임에 조건에 따라 다른 실행 경로를 선택하는 동적 분기 로직을 구현하는 것이 유용합니다.
4. **리뷰의 중요성:** `sg_id` 클램핑과 같은 OOB 접근 방지 로직은 코드 리뷰를 통해 발견되고 수정되었습니다. 이는 복잡한 병렬 프로그래밍에서 안정성과 정확성을 보장하는 데 있어 코드 리뷰의 중요성을 강조합니다.
## References
* [WGSL Specification - Invalid Load](https://www.w3.org/TR/WGSL/#invalid-load)
* [WebGPU API](https://www.w3.org/TR/webgpu/)
* [ONNX Runtime](https://onnxruntime.ai/)
* [MatMulNBits Wide Tile Kernel Logic](https://github.com/microsoft/onnxruntime/blob/main/onnxruntime/contrib_ops/webgpu/quantization/matmul_nbits_wide_tile.wgsl.template)
## 참고 자료
- https://www.w3.org/TR/WGSL/#invalid-load
- https://www.w3.org/TR/webgpu/
- https://onnxruntime.ai/
- https://github.com/microsoft/onnxruntime/blob/main/onnxruntime/contrib_ops/webgpu/quantization/matmul_nbits_wide_tile.wgsl.template
> ⚠️ **알림:** 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [onnxruntime] ONNX Runtime WebGPU EP: 디바이스 없는 오프라인 컴파일 지원
- [onnxruntime] ONNX Runtime: Arm64 KleidiAI 기반 FP16 GEMM 및 Convolution 최적화
- [onnxruntime] ONNX Runtime: AVX2 및 AVX-VNNI를 위한 2-bit 가중치 CPU 커널 최적화
- [onnxruntime] ONNX Runtime WebGPU: Intel Xe-3LPG를 위한 고성능 GEMM 최적화 분석
- [onnxruntime] ONNX Runtime: MoE Router GEMV 최적화 및 Bias Fusion 구현
PR Analysis 의 다른글
- 이전글 [sglang] SGLang의 FLUX.2 추론 최적화: Eager 모드에서의 커널 융합 기법
- 현재글 : [onnxruntime] WebGPU MatMulNBits 최적화: Subgroup Shuffle을 활용한 성능 향상
- 다음글 [sglang] SGLang의 Wan2.2-TI2V 최적화: Triton 커널을 통한 메모리 트래픽 병목 해결
댓글