Skip to content

About

GPU Kernel Optimisation for GLM-ASR Inference via Triton

Resources

Stars

0 stars

Watchers

0 watching

Forks

Latest commit

 

History

4 Commits

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

⚡ GPU Kernel Optimisation for GLM-ASR Inference via Triton

An end-to-end performance engineering framework achieving a 4.26× total acceleration on multimodal speech-to-text inference. Developed as part of the University of Edinburgh's Machine Learning Systems curriculum (INFR11269), this system compresses total decoding latency from 1493.1 ms down to 350.2 ms (dropping per-token latency from 114.85 ms to 26.94 ms) on an NVIDIA H200 Tensor Core GPU (141GB HBM3) with zero accuracy loss.


❌ The Hardware Constraint & Bottleneck Analysis

Autoregressive language model decoding at a batch size of 1 is fundamentally memory-bandwidth bound, not compute-bound. The GPU's massive Tensor Core throughput is left idle while the 3.35 TB/s High-Bandwidth Memory (HBM3) bus acts as the binding constraint.

The Baseline Overhead Profile:

  1. Redundant VRAM Traffic: The unoptimized baseline constantly loads float32 weights and streams intermediate attention score matrices back and forth to global memory.
  2. Launch Latency Accumulation: PyTorch eager mode emits 392 mini-kernel dispatches per decode step for Rotary Position Embeddings (RoPE) alone. At ~5–10 µs per CUDA launch, dispatch overhead costs ~3 ms per token step before any computation occurs.
  3. Dynamic Memory Allocation: Frequent torch.cat operations during KV-cache extension cause 1,400 transient allocations across a typical 50-step generation sequence.

🏗️ System Architecture & Optimization Matrix

This framework overrides baseline abstractions with hand-rolled Triton kernels, keeping computational tiles local to fast on-chip registers and SRAM.

graph TD
    %% Baseline Path
    subgraph Baseline Eager Path [~476 CUDA Dispatches / Layer Token Write back]
        A[Hidden States] --> B[Linear Projections]
        B -->|Write F32 to HBM3| C[Global VRAM]
        C -->|Launch RoPE Mini-Kernels| D[PyTorch RoPE Fallback]
        D -->|Write Intermediate Scores| E[3-Kernel Attention Loop]
        E -->|Dynamic torch.cat Allocation| F[Growing KV-Cache Tensor]
    end

    %% Optimized Triton Path
    subgraph Optimized Triton Path [~56 Fused Compute Steps / In-Register Execution]
        G[Hidden States] --> H[Unconditional Triton Norms]
        H -->|Direct bf16 Weight Caching| I[cuBLAS HGEMM Tensor Core Path]
        I -->|Fused Single-Store Register Shuffling| J[fused_rope_apply_kernel]
        J -->|Single-Pass Safe Softmax SRAM Tiling| K[fused_attention_decode_kernel]
        K -->|Pre-Allocated Static In-Place Updates| L[Fixed In-Place KV Buffer]
    end

    %% Contrast Color Accents
    style C fill:#ff9999,stroke:#b30000,stroke-width:2px
    style E fill:#ff9999,stroke:#b30000,stroke-width:2px
    style J fill:#99ff99,stroke:#006600,stroke-width:2px
    style K fill:#99ff99,stroke:#006600,stroke-width:2px
Loading

Core Low-Level Engineering Vectors:

  • Persistent Weight Caching & Precision Alignment (best_ensemble_v2): Caches model parameters as persistent bfloat16 attributes upon the initial forward pass. This cuts memory bandwidth requirements in half (2 B vs 4 B per element) and forces execution down the vendor-tuned cuBLAS HGEMM Tensor Core path.
  • FlashAttention-Style Single-Pass Decode (best_ensemble_v3): Implements a single fused_attention_decode_kernel for sequence lengths $\leq 512$. Because the whole sequence fits inside a single SRAM tile (BLOCK_K=next_power_of_two(seq_k), BLOCK_D=128), it executes an efficient, max-stable safe softmax entirely in registers, bypassing any outer loop or global memory writes for intermediate attention scores.
  • Structural RoPE Fusion (best_ensemble_v4): Consolidates rotation operations into a single execution pass. It collapses dimensions into a (batch × heads, seq_position) grid mapping, squeezing out 336 redundant kernel launches per step.
  • Static In-Place Buffering: Replaces dynamic memory allocation with a dedicated forward_with_kv_buffer() path that updates pre-allocated, fixed-size matrices in place.

📈 Roofline Model & Operator Classification

By mapping metrics against the NVIDIA H200 Ridge Point ($\sim1,181\text{ FLOP/byte}$), the system selectively targets optimization where memory-bandwidth constraints are highest.

Active Triton Kernel Operational Tensor Shape Mathematical Flop Volume Arithmetic Intensity (AI) Binding Hardware Constraint
rmsnorm_kernel $(1, 2048)$ $\sim 4\text{ K}$ $\sim 1\text{ FLOP/B}$ Memory Bandwidth Bound
silu_kernel / gelu_kernel $(1, N)$ $\sim 5N$ $\sim 0.6\text{ FLOP/B}$ Memory Bandwidth Bound
fused_rope_apply_kernel $(28, 1, 128)$ $\sim 14\text{ K}$ $\sim 2\text{ FLOP/B}$ Memory Bandwidth Bound
fused_attention_decode_kernel $(16, 1, 100)$ $\sim 410\text{ K}$ $\sim 3\text{ FLOP/B}$ Memory Bandwidth Bound
Linear (Decoder GEMV) $(1, 2048, 2048)$ $8\text{ M}$ $\sim 1\text{ FLOP/B}$ Memory Bandwidth Bound
Linear (Encoder Prefill GEMM) $(750, 2048, 2048)$ $6.3\text{ G}$ $\sim 214\text{ FLOP/B}$ Mixed / Ridge Boundary
Linear (Prefill MLP Step) $(750, 2048, 5632)$ $17.3\text{ G}$ $\sim 586\text{ FLOP/B}$ Compute Bound

📊 Performance Iteration Milestones

Incremental evaluation sweeps demonstrate how algorithmic choices and precision strategy dominate over hyperparameter tuning:

Cumulative End-to-End Latency Profile

Optimization Version Milestone Global Latency (ms) Step Velocity (ms/token) Net Speedup Factor Primary Architectural Modification
triton_example (Baseline) 1493.1 ms 114.85 ms/tok $1.00\times$ Reference implementation.
best_ensemble (v1) 1114.0 ms 85.70 ms/tok $1.34\times$ Routed 750-frame prefill via FlashAttention-2.
best_ensemble_v2 486.0 ms 37.40 ms/tok $3.07\times$ bfloat16 weight caching; activated H200 Tensor Cores.
best_ensemble_v3 413.0 ms 31.80 ms/tok $3.62\times$ Injected fused single-pass decode attention.
best_ensemble_v4 (Final Submission) 350.2 ms 26.94 ms/tok $4.26\times$ MAX_DIM 512 expansion + structural RoPE fusion.

Component Operator Breakdown

Operational Layer Profile Baseline Metric Optimized Metric (v4) Local Layer Speedup
Audio Encoder (32 Layers) 1035.58 ms 523.37 ms 1.98× Speedup
Text Decoder Layer 0 1.69 ms 0.89 ms 1.90× Speedup
Full Decoder MLP (SwiGLU Block) 50.78 ms 5.04 ms 10.08× Speedup
Decode Loop Attention Step 0.61 ms 0.33 ms 1.85× Speedup

📁 Repository Blueprint & Modular Architecture

gpu-kernel-optimisation-glm-asr/
├── README.md               # Hardware architectural specifications and benchmarking data
├── GUIDE.md                # Technical walkthrough outlining manual pointer layout modifications
├── requirements.txt       # Pinlocked execution footprints ensuring environment consistency
├── benchmark.sh            # Ingests model presets to profile steady-state latency sweeps
├── benchmark_detailed.sh   # Sub-component profiling utility leveraging raw CUDA event timing
├── demo.py                 # Core orchestration runtime processing raw test waveform inputs
├── model.py                # Model architecture overriding standard layers with Triton backends
├── layers.py               # Custom operational layers managing memory caching strategies
├── attention.py            # FlashAttention-inspired single-pass decoding logic
└── audio_sample.wav        # 3.50s 16kHz verification audio source

🚀 Environment Execution Guide

Infrastructure Hardware Prerequisites

  • Environment Node: Configured for high-performance clusters (e.g., UoE SAXA Node allocated via SLURM).
  • Hardware Baseline: Ampere, Hopper, or Blackwell GPU with Compute Capability 8.0+ (Optimized for NVIDIA H200 SXM5).
  • Software Base: Linux OS, CUDA 12.x, Python 3.12+, PyTorch 2.x.

Running the System

# Clone the optimized framework repository
git clone https://github.com
cd gpu-kernel-optimisation-glm-asr

# Install the required software stack
pip install -r requirements.txt

# Run full end-to-end ASR audio transcription verification
python demo.py

# Run the target performance latency evaluation sweeps
bash benchmark.sh

📚 References

About

GPU Kernel Optimisation for GLM-ASR Inference via Triton

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages