GitHub - PJHkorea/pim-hbm-bypass: Software-Defined PIM-HBM Bypass Infrastructure. 0ns JAX/XLA memory view fusion with real-time asynchronous NCCL collective fault-scanning & no-recompile Mux hot-swapping.

6 min read Original article โ†—

๐Ÿš€ pim-hbm-bypass Blueprint

Experimental PIM-HBM Hardware Co-Design Subsystem exploring 0ns Framework Memory View Fusion & Fault Telemetry

This project is a hardware-software co-design prototype engineered to investigate and resolve software abstraction fragmentation barriers within next-generation accelerator infrastructure environments.

By interlocking the low-level physical cache-line alignment mechanisms directly with the upper high-performance framework (JAX/XLA) address interfaces into a unified algebraic pipeline, this subsystem explores the structural viability of stably mitigating static compilation stalls and hardware timing jitters during distributed cluster operations.


โšก Key Innovations

  1. 0ns Memory Copy Bypass: Leverages the __cuda_array_interface__ specification to directly link low-level C++ physical address lines to the JAX tensor bus, structurally eliminating host-device hardware memory copy overhead.
  2. Pure Branchless Loop: Completely eradicates conditional branches (if), instead deploying precise ternary operations and register-resident data reuse patterns to guarantee the compiler forces conditional move primitives (SEL/PRMT).
  3. Warp-level Dynamic Bounds: Intercepts potential out-of-bound memory faults (SegFault) inside ragged tail grids by driving a warp-level binary tree reduction firewall via __shfl_down_sync to compute active maximum surviving address offsets.
  4. Algebraic Insulation Gate: Captures hardware fault signals and floating-point divergence anomalies behind the JAX runtime using dedicated stop_gradient circuits, isolating backpropagation chain contamination at the physical layer.
  5. Dynamic Hot-Plugging Recovery: Upon physical HBM bank failure during active cluster runs, this mechanism attempts real-time 64-bit address wire hot-swapping for the corrupted device slot without causing collective communication (NCCL) stalls or graph recompilation.

๐Ÿ“‚ Repository Structure

  • LICENSE: Declares legal safeguards and patent retaliation defense clauses under the Apache License 2.0 specification.
  • CMakeLists.txt: The build orchestrator that automatically tracks system hardware architecture topologies and pybind11 compilation paths to emit the final shared object (.so) libraries.
  • pim_hbm_core.cu: The core branchless mathematical acceleration kernel implementing alignas(32) cache-line matching, __activemask() dynamic address firewalls, and __ldg high-speed read rails (C++/CUDA).
  • pim_hardware_gate.py: The pre-warming and backpropagation chain insulation layer utilizing ShapeDtypeStruct virtual abstract tracers to lock XLA compiler machine code while maintaining 0MB of physical VRAM footprint (Python/JAX).
  • topology_sharding.py: The macro-level topology control tower that intercepts per-node VRAM physical address lines to establish zero-copy NamedSharding global distributed matrix views (Python/JAX).
  • hardware_fault_recovery.py: The real-time, high-availability hot-plugging swap engine that handles background fault scans within the distributed weight matrices and switches routing to the pre-reserved emergency backup address pool (Python/JAX).
  • hardware_fault_recovery_distributed.py: The collective scanning and recovery engine engineered for large-scale infrastructure, leveraging wire-level NCCL All-Reduce fusion and np.flatnonzero vectorized fault extraction (Python/JAX).
  • llama3_layer_adapter.py: A transformer layer adapter plugin that conducts direct 0ns address ingestion matching Llama-3-8B dimensions (4096 / 14336) alongside branchless fault-tolerant forward execution buses (Python/JAX).

๐Ÿ› ๏ธ Quick Start

Execute the following commands sequentially within a high-performance cluster terminal configured with NVIDIA Ampere (A100) or Hopper (H100/H200) environments to ignite the memory view fusion mode and the hardware fault-tolerant emulation engine.

# 1. Create and enter an isolated directory dedicated to the build sequence
mkdir build && cd build

# 2. Launch cross-compilation architecture scanning (Automatically tracks pybind11 and CUDA paths)
cmake ..

# 3. Build and compile the hardware machine-code library (Extracts pim_hbm_bridge_core.so)
make -j\$(nproc)

# 4. Migrate the extracted shared library module back to the upper execution directory
cp pim_hbm_bridge_core*.so .. && cd ..

# 5. [STEP A] Execute the single-node JAX pre-warming engine and algebraic insulation guard
python3 pim_hardware_gate.py

# 6. [STEP B] Launch the macro-level distributed sharding topology virtual view fusion matrix
python3 topology_sharding.py

# 7. [STEP C] Run the collective health scan and dynamic hot-plugging recovery via wire-level NCCL All-Reduce
python3 hardware_fault_recovery_distributed.py

# 8. [โšก STEP D] Instantiate the zero-copy address adapter infrastructure tailored for Llama-3-8B (4096/14336) matrices
python3 llama3_layer_adapter.py

๐Ÿ“œ License

This project is distributed under the terms of the Apache License 2.0. You are free to modify and distribute the software, provided that the original copyright notice and license disclosure obligations are fully preserved.


๐Ÿš€ pim-hbm-bypass Blueprint

Experimental PIM-HBM Hardware Co-Design Subsystem exploring 0ns Framework Memory View Fusion & Fault Telemetry

๋ณธ ํ”„๋กœ์ ํŠธ๋Š” ์ฐจ์„ธ๋Œ€ ๊ฐ€์†๊ธฐ ์ธํ”„๋ผ ํ™˜๊ฒฝ์—์„œ ๋ฐœ์ƒํ•  ์ˆ˜ ์žˆ๋Š” ์†Œํ”„ํŠธ์›จ์–ด ๊ณ„์ธต ๊ฐ„ ํŒŒํŽธํ™” ์žฅ๋ฒฝ์„ ์—ฐ๊ตฌํ•˜๊ธฐ ์œ„ํ•ด ์„ค๊ณ„๋œ ํ•˜๋“œ์›จ์–ด-์†Œํ”„ํŠธ์›จ์–ด ๊ณต๋™ ์„ค๊ณ„(Co-design) ํ”„๋กœํ† ํƒ€์ž…์ž…๋‹ˆ๋‹ค.

์ €์ˆ˜์ค€์˜ ๋ฌผ๋ฆฌ ์บ์‹œ๋ผ์ธ ์ •๋ ฌ ๋งค์ปค๋‹ˆ์ฆ˜๊ณผ ์ƒ์œ„ ๊ณ ์„ฑ๋Šฅ ํ”„๋ ˆ์ž„์›Œํฌ(JAX/XLA) ๊ฐ„์˜ ์ฃผ์†Œ์„  ์ธํ„ฐํŽ˜์ด์Šค๋ฅผ ๋‹จ์ผ ๋Œ€์ˆ˜ ํŒŒ์ดํ”„๋ผ์ธ์œผ๋กœ ์—ฐ๊ฒฐํ•˜์—ฌ, ๋ถ„์‚ฐ ํด๋Ÿฌ์Šคํ„ฐ ๊ฐ€๋™ ์ค‘ ์œ ๋ฐœ๋˜๋Š” ์ •์  ์ปดํŒŒ์ผ ๋ ‰ ๋ฐ ํ•˜๋“œ์›จ์–ด ์ง€ํ„ฐ๋ฅผ ์•ˆ์ •์ ์œผ๋กœ ์ œ์–ดํ•  ์ˆ˜ ์žˆ๋Š” ๊ฐ€๋Šฅ์„ฑ์„ ํƒ๊ตฌํ•ฉ๋‹ˆ๋‹ค.


โšก ํ•ต์‹ฌ ์•„ํ‚คํ…์ฒ˜ ํŠน์„ฑ (Key Innovations)

  1. 0ns Memory Copy Bypass: __cuda_array_interface__ ๊ทœ๊ฒฉ์„ ํ™œ์šฉํ•ด C++ ๊ธฐ๊ณ„์–ด ์ฃผ์†Œ์„ ์„ JAX ํ…์„œ ๋ฒ„์Šค์— ์ง๊ฒฐํ•จ์œผ๋กœ์จ, ํ˜ธ์ŠคํŠธ-๋””๋ฐ”์ด์Šค ๊ฐ„ ๋ฌผ๋ฆฌ ๋ณต์‚ฌ ์˜ค๋ฒ„ํ—ค๋“œ๋ฅผ ๊ตฌ์กฐ์ ์œผ๋กœ ํ•ด์†Œํ•ฉ๋‹ˆ๋‹ค.
  2. Pure Branchless Loop: ์กฐ๊ฑด ๋ถ„๊ธฐ๋ฌธ(if)์„ ์™„์ „ํžˆ ๋ฐฐ์ œํ•˜๊ณ  ์‚ผํ•ญ ์—ฐ์‚ฐ ๋ฐ ๋ ˆ์ง€์Šคํ„ฐ ์ƒ์ฃผ ๋ฐ์ดํ„ฐ ์žฌ์‚ฌ์šฉ ๊ตฌ์กฐ๋ฅผ ์ •๋ฐ€ ๊ตฌ์„ฑํ•˜์—ฌ, ์ปดํŒŒ์ผ๋Ÿฌ ์ˆ˜์ค€์˜ ์กฐ๊ฑด๋ถ€ ์ด๋™ ๋ช…๋ น์–ด(SEL/PRMT) ์ถœ๋ ฅ์„ ์œ ๋„ํ•ฉ๋‹ˆ๋‹ค.
  3. Warp-level Dynamic Bounds: ๋งˆ์ง€๋ง‰ ๊ทธ๋ฆฌ๋“œ ์žํˆฌ๋ฆฌ ์˜์—ญ์—์„œ ๋ฐœ์ƒํ•  ์ˆ˜ ์žˆ๋Š” ๋ฉ”๋ชจ๋ฆฌ ์ฐธ์กฐ ์˜ค๋ฅ˜(SegFault)๋ฅผ ๋ฐฉ์ง€ํ•˜๊ธฐ ์œ„ํ•ด, __shfl_down_sync ๊ธฐ๋ฐ˜ ์›Œํ”„ ๋‚ด 2์ง„ ํŠธ๋ฆฌ ์ตœ๋Œ€ ์ƒ์กด ์ฃผ์†Œ ๋™์  ๋ฆฌ๋•์…˜ ๋ฐฉํ™”๋ฒฝ์„ ๊ตฌ๋™ํ•ฉ๋‹ˆ๋‹ค.
  4. Algebraic Insulation Gate: JAX ๋Ÿฐํƒ€์ž„์—์„œ ํ•˜๋“œ์›จ์–ด ๊ฒฐํ•จ ์‹ ํ˜ธ ๋ฐ ์ˆ˜์น˜ ๋ฐœ์‚ฐ ์˜ค์ฐจ๋ฅผ stop_gradient ํšŒ๋กœ๋กœ ํฌํšํ•˜์—ฌ, ๋ฏธ๋ถ„ ์‚ฌ์Šฌ ์˜ค์—ผ์„ ํ”ผ์ง€์ปฌ ๋ ˆ๋ฒจ์—์„œ ๊ฒฉ๋ฆฌ ๋ฐ ์ ˆ์—ฐํ•ฉ๋‹ˆ๋‹ค.
  5. Dynamic Hot-Plugging Recovery: ๊ฐ€์†๊ธฐ ํด๋Ÿฌ์Šคํ„ฐ ๊ตฌ๋™ ์ค‘ ํŠน์ • HBM ๋ฑ…ํฌ์— ๋ฌผ๋ฆฌ์  ๊ฒฐํ•จ ๋ฐœ์ƒ ์‹œ, ๊ธ€๋กœ๋ฒŒ ํ†ต์‹ (NCCL) ์ค‘๋‹จ ๋ฐ ๊ทธ๋ž˜ํ”„ ์žฌ์ปดํŒŒ์ผ ์—†์ด ์˜ค์ง ๋ถˆ๋Ÿ‰ ์žฅ์น˜์˜ 64๋น„ํŠธ ์ฃผ์†Œ์„ ๋งŒ ์‹ค์‹œ๊ฐ„์œผ๋กœ ์šฐํšŒ ์Šค์™€ํ•‘(Hot-Swapping)ํ•˜๋Š” ๋ฉ”์ปค๋‹ˆ์ฆ˜์„ ์‹œ๋„ํ•ฉ๋‹ˆ๋‹ค.

๐Ÿ“‚ ํŒŒ์ผ ํ† ํด๋กœ์ง€ (Repository Structure)

  • LICENSE: Apache License 2.0 ์˜๊ฑฐ ๋ฒ•์  ๋ฐฉํ™”๋ฒฝ ๋ฐ ํŠนํ—ˆ ๋ณดํ˜ธ ์กฐํ•ญ ๋ช…์‹œ
  • CMakeLists.txt: ์‹œ์Šคํ…œ ๊ฐ€์†๊ธฐ ์•„ํ‚คํ…์ฒ˜ ํ™˜๊ฒฝ๊ณผ pybind11 ์ปดํŒŒ์ผ ํŒจ์Šค๋ฅผ ์ž๋™ ์ถ”์ ํ•˜์—ฌ ๊ณต์œ  ๋ผ์ด๋ธŒ๋Ÿฌ๋ฆฌ(.so)๋ฅผ ์‚ฌ์ถœํ•˜๋Š” ๋นŒ๋“œ ์˜ค์ผ€์ŠคํŠธ๋ ˆ์ดํ„ฐ
  • pim_hbm_core.cu: alignas(32) ์บ์‹œ๋ผ์ธ ์ผ์น˜ ๋ ˆ์ด์•„์›ƒ, __activemask() ๋™์  ์ฃผ์†Œ ๋ฐฉํ™”๋ฒฝ ๋ฐ __ldg ๊ฐ€์† ๋ ˆ์ผ์ด ์ฃผ์ž…๋œ ๋ฌด๋ถ„๊ธฐ ์ˆ˜ํ•™ ๊ฐ€์† ์ปค๋„ ์ฝ”์–ด (C++/CUDA)
  • pim_hardware_gate.py: ShapeDtypeStruct ๊ฐ€์ƒ ์ถ”์ƒํ™” ํŠธ๋ ˆ์ด์„œ๋ฅผ ํ™œ์šฉํ•ด ์‹ค์žฌ VRAM ์ ์œ  0MB ์ƒํƒœ๋กœ XLA ์ปดํŒŒ์ผ๋Ÿฌ ๊ธฐ๊ณ„์–ด๋ฅผ ๊ณ ์ •ํ•˜๋Š” ์˜ˆ์—ด ๋ฐ ๋ฏธ๋ถ„ ์‚ฌ์Šฌ ์ ˆ์—ฐ ๋ ˆ์ด์–ด (Python/JAX)
  • topology_sharding.py: ๋Œ€๊ทœ๋ชจ ํด๋Ÿฌ์Šคํ„ฐ ๋…ธ๋“œ๋ณ„ VRAM ๋ฌผ๋ฆฌ ์ฃผ์†Œ์„ ์„ ๊ฐ€๋กœ์ฑ„์–ด ์ œ๋กœ์นดํ”ผ NamedSharding ๊ธ€๋กœ๋ฒŒ ๋ถ„์‚ฐ ๋งคํŠธ๋ฆญ์Šค ๋ทฐ๋ฅผ ์ˆ˜๋ฆฝํ•˜๋Š” ๊ฑฐ์‹œ์  ํ† ํด๋กœ์ง€ ๊ด€์ œํƒ‘ (Python/JAX)
  • hardware_fault_recovery.py: ๋ถ„์‚ฐ ๊ฐ€์ค‘์น˜ ํ–‰๋ ฌ ๋‚ด ๋ถˆ๋Ÿ‰ ๋ฑ…ํฌ ๋ฐฑ๊ทธ๋ผ์šด๋“œ ์Šค์บ” ๋ฐ ๋น„์ƒ ๋ฐฑ์—… ํ’€ ์ฃผ์†Œ์„ ์„ ํ™œ์šฉํ•œ ์‹ค์‹œ๊ฐ„ ๋ฌด์ค‘๋‹จ ํ•ซํ”Œ๋Ÿฌ๊น… ์Šค์™€ํ”„ ์—”์ง„ (Python/JAX)
  • hardware_fault_recovery_distributed.py: ์ดˆ๋Œ€ํ˜• ์ธํ”„๋ผ๋ฅผ ์œ„ํ•œ NCCL All-Reduce ์™€์ด์–ด ๋ ˆ๋ฒจ ์œตํ•ฉ ์ง‘์‚ฐ(Collective) ์Šค์บ” ๋ฐ np.flatnonzero ๋ฒกํ„ฐํ™” ๊ฒฐํ•จ ์ ์ถœ ๋ณต๊ตฌ ์—”์ง„ (Python/JAX)
  • llama3_layer_adapter.py: Llama-3-8B ๊ณ ์œ  ์ฐจ์›(4096 / 14336) ์งํ†ต 0ns ์ฃผ์†Œ ์ธ์ž… ๋ฐ ๋ฌด๋ถ„๊ธฐ ๊ฒฐํ•จ ํ—ˆ์šฉ ์ˆœ๋ฐฉํ–ฅ ํ›ˆ๋ จ/์ถ”๋ก  ๋ฒ„์Šค ์–ด๋Œ‘ํ„ฐ ํ”Œ๋Ÿฌ๊ทธ์ธ (Python/JAX)

๐Ÿ› ๏ธ ๊ณ ์† ๊ตฌ๋™ ๋ฐ ๋นŒ๋“œ ์ง€์นจ (Quick Start)

NVIDIA Ampere(A100) ๋˜๋Š” Hopper(H100/H200) ํ™˜๊ฒฝ์ด ๊ตฌ์ถ•๋œ ๊ณ ์„ฑ๋Šฅ ํด๋Ÿฌ์Šคํ„ฐ ํ„ฐ๋ฏธ๋„์—์„œ ๋‹ค์Œ ๋ช…๋ น์„ ์ˆœ์ฐจ ๊ฐ€๋™ํ•˜์—ฌ ๋ฐ”์ดํŒจ์Šค ๋ชจ๋“œ ๋ฐ ํ•˜๋“œ์›จ์–ด ๊ฒฐํ•จ ํ—ˆ์šฉ(Fault-Tolerant) ์—”์ง„ ์—๋ฎฌ๋ ˆ์ด์…˜์„ ๊ธฐํญํ•  ์ˆ˜ ์žˆ์Šต๋‹ˆ๋‹ค.

# 1. ๋นŒ๋“œ ์ „์šฉ ๊ฒฉ๋ฆฌ ๊ณต๊ฐ„ ์ƒ์„ฑ ๋ฐ ์ง„์ž…
mkdir build && cd build

# 2. ํฌ๋กœ์Šค ์ปดํŒŒ์ผ ์•„ํ‚คํ…์ฒ˜ ์Šค์บ” ๊ฐ€๋™ (pybind11 ๋ฐ CUDA ์ปดํŒŒ์ผ ํŒจ์Šค ์ž๋™ ์ถ”์ )
cmake ..

# 3. ํ•˜๋“œ์›จ์–ด ๊ธฐ๊ณ„์–ด ๋ผ์ด๋ธŒ๋Ÿฌ๋ฆฌ ์ปดํŒŒ์ผ ๋นŒ๋“œ (pim_hbm_bridge_core.so ์ถ”์ถœ)
make -j\$(nproc)

# 4. ์ƒ์œ„ ์‹คํ–‰ ๋””๋ ‰ํ† ๋ฆฌ๋กœ ์ถ”์ถœ๋œ ๋ชจ๋“ˆ ๋ณต์‚ฌ ์ด๊ด€
cp pim_hbm_bridge_core*.so .. && cd ..

# 5. [STEP A] ๋‹จ์ผ ๋…ธ๋“œ ์ „์šฉ JAX ์˜ˆ์—ด ์—”์ง„ ๋ฐ ๋Œ€์ˆ˜์  ์˜ค์ฐจ ์ ˆ์—ฐ ๊ฐ€๋“œ ๊ฐ€๋™
python3 pim_hardware_gate.py

# 6. [STEP B] Multi-GPU ๊ฑฐ์‹œ ๋ถ„์‚ฐ ์ƒค๋”ฉ ํ† ํด๋กœ์ง€ ๊ฐ€์ƒ ์œตํ•ฉ ํ…์„œ ๊ฐ€๋™ 
python3 topology_sharding.py

# 7. [STEP C] ์ดˆ๋Œ€ํ˜• ์ธํ”„๋ผ์šฉ NCCL All-Reduce ์œตํ•ฉ ๋ถ„์‚ฐ ์ง‘์‚ฐ ํ—ฌ์Šค ์Šค์บ” ๋ฐ ํ•ซํ”Œ๋Ÿฌ๊น… ๋ณต๊ตฌ ๊ฐ€๋™
python3 hardware_fault_recovery_distributed.py

# 8. [โšก STEP D] ์‹ค์ „ Llama-3-8B Transformer 4096/14336 ๋งคํŠธ๋ฆญ์Šค 0ns ์ฃผ์†Œ ์ œ๋กœ์นดํ”ผ ์–ด๋Œ‘ํ„ฐ ์ธํ”„๋ผ ๊ฐ€๋™
python3 llama3_layer_adapter.py

๐Ÿ“œ ๋ผ์ด์„ ์Šค (License)

๋ณธ ํ”„๋กœ์ ํŠธ๋Š” Apache License 2.0 ์˜๊ฑฐํ•˜์—ฌ ๋ฐฐํฌ๋ฉ๋‹ˆ๋‹ค. ์ž์œ ๋กœ์šด ์ˆ˜์ • ๋ฐ ๋ฐฐํฌ๊ฐ€ ๊ฐ€๋Šฅํ•˜๋‚˜ ์ €์ž‘๊ถŒ ๋ฐ ๋ผ์ด์„ ์Šค ๊ณ ์ง€ ์˜๋ฌด๊ฐ€ ์ˆ˜๋ฐ˜๋ฉ๋‹ˆ๋‹ค.