๐Ÿง  Accelerator

XLA ์ƒ์„ธ

๊ฐœ์š”

XLA(Accelerated Linear Algebra)๋Š” Google์ด ๊ฐœ๋ฐœํ•œ ์˜คํ”ˆ์†Œ์Šค ๋จธ์‹ ๋Ÿฌ๋‹ ์ปดํŒŒ์ผ๋Ÿฌ ํ”„๋ ˆ์ž„์›Œํฌ๋กœ, TensorFlow ๋ฐ JAX ๋ชจ๋ธ์„ ๋‹ค์–‘ํ•œ ํ•˜๋“œ์›จ์–ด ๊ฐ€์†๊ธฐ์—์„œ ํšจ์œจ์ ์œผ๋กœ ์‹คํ–‰ํ•˜๊ธฐ ์œ„ํ•ด ์ตœ์ ํ™”๋œ ์ค‘๊ฐ„ ํ‘œํ˜„(IR)๊ณผ ์ฝ”๋“œ ์ƒ์„ฑ ๊ธฐ๋Šฅ์„ ์ œ๊ณตํ•œ๋‹ค. 2017๋…„์— ์†Œ๊ฐœ๋œ ์ดํ›„ TPU, GPU, CPU ๋“ฑ ๊ด‘๋ฒ”์œ„ํ•œ ํ•˜๋“œ์›จ์–ด๋ฅผ ์ง€์›ํ•˜๋ฉฐ, ์ตœ๊ทผ์—๋Š” StableHLO๋ฅผ ํ†ตํ•ด ํ”„๋ ˆ์ž„์›Œํฌ ๋…๋ฆฝ์ ์ธ ํ˜ธํ™˜์„ฑ์„ ์ถ”๊ตฌํ•˜๊ณ  ์žˆ๋‹ค.

XLA์˜ ํ•ต์‹ฌ ์„ค๊ณ„ ์ฒ ํ•™์€ "์ง€์—ฐ ์ปดํŒŒ์ผ(lazy compilation)"๊ณผ "์—”๋“œํˆฌ์—”๋“œ ์ตœ์ ํ™”"์ด๋‹ค. ๋ชจ๋ธ์˜ ์—ฐ์‚ฐ ๊ทธ๋ž˜ํ”„๋ฅผ HLO(High Level Optimizer) ์ค‘๊ฐ„ ํ‘œํ˜„์œผ๋กœ ๋ณ€ํ™˜ํ•œ ํ›„, ํ•˜๋“œ์›จ์–ด ํŠนํ™” ์ตœ์ ํ™”์™€ ์ฝ”๋“œ ์ƒ์„ฑ์„ ์ˆ˜ํ–‰ํ•œ๋‹ค. ์ด๋ฅผ ํ†ตํ•ด ๋ฉ”๋ชจ๋ฆฌ ์ ‘๊ทผ ํŒจํ„ด ์ตœ์ ํ™”, ์—ฐ์‚ฐ ์œตํ•ฉ, ๋ณ‘๋ ฌํ™” ๋“ฑ์„ ์ž๋™์œผ๋กœ ์ˆ˜ํ–‰ํ•˜์—ฌ ๋†’์€ ์„ฑ๋Šฅ์„ ๋‹ฌ์„ฑํ•œ๋‹ค.

ํ•ต์‹ฌ ๊ฐœ๋…

XLA ์Šคํƒ ์•„ํ‚คํ…์ฒ˜

XLA Stack Architecture

XLA ์Šคํƒ์€ ํฌ๊ฒŒ ๋„ค ๊ฐ€์ง€ ๊ณ„์ธต์œผ๋กœ ๊ตฌ์„ฑ๋œ๋‹ค:

๊ณ„์ธต ๋ชจ๋“ˆ ์„ค๋ช…
ํ”„๋ก ํŠธ์—”๋“œ TensorFlow/JAX ๋ชจ๋ธ์„ HLO๋กœ ๋ณ€ํ™˜
์ค‘๊ฐ„ ํ‘œํ˜„ HLO (High Level Optimizer) ๊ทธ๋ž˜ํ”„ ์ˆ˜์ค€ ์ถ”์ƒํ™”
์ตœ์ ํ™” XLA Compiler ์—ฐ์‚ฐ ์œตํ•ฉ, ๋ฉ”๋ชจ๋ฆฌ ์ตœ์ ํ™”
๋ฐฑ์—”๋“œ Target Specific TPU/GPU/CPU๋ณ„ ์ฝ”๋“œ ์ƒ์„ฑ

HLO: ํ•ต์‹ฌ ์ค‘๊ฐ„ ํ‘œํ˜„

XLA์˜ ๋ชจ๋“  ์ตœ์ ํ™”๋Š” HLO(High Level Optimizer) ์ค‘๊ฐ„ ํ‘œํ˜„์„ ์ค‘์‹ฌ์œผ๋กœ ์ด๋ฃจ์–ด์ง„๋‹ค. HLO๋Š” ํ”„๋ ˆ์ž„์›Œํฌ ๋…๋ฆฝ์ ์ธ ์ •๋ ฌ๋œ ์–ด์…ˆ๋ธ”๋ฆฌ์™€ ๊ฐ™์€ ์—ญํ• ์„ ํ•˜๋ฉฐ, ๋‹ค์Œ ํŠน์ง•์„ ๊ฐ€์ง„๋‹ค:

  • HLO Instruction: ๊ฐœ๋ณ„ ์—ฐ์‚ฐ์„ ํ‘œํ˜„ํ•˜๋Š” ๊ธฐ๋ณธ ๋‹จ์œ„. add, multiply, convolution ๋“ฑ์˜ ๋ช…๋ น์–ด๋กœ ๊ตฌ์„ฑ๋œ๋‹ค.
  • HLO Computation: ๋ช…๋ น์–ด๋“ค์˜ ๊ทธ๋ž˜ํ”„ ๊ตฌ์กฐ. ๋ชจ๋ธ์˜ ์„œ๋ธŒ๊ทธ๋ž˜ํ”„์— ๋Œ€์‘ํ•œ๋‹ค.
  • HLO Module: ์ „์ฒด ๋ชจ๋ธ์„ ๋‚˜ํƒ€๋‚ด๋Š” ์ตœ์ƒ์œ„ ๊ตฌ์กฐ. ์—ฌ๋Ÿฌ Computation์„ ํฌํ•จํ•œ๋‹ค.

StableHLO: ํ”„๋ ˆ์ž„์›Œํฌ ๋…๋ฆฝ์  IR

StableHLO Overview

StableHLO๋Š” XLA์˜ ๋ฐœ์ „๋œ ํ˜•ํƒœ๋กœ, ํ”„๋ ˆ์ž„์›Œํฌ ๊ฐ„ ํ˜ธํ™˜์„ฑ์„ ์ถ”๊ตฌํ•˜๋Š” ์˜คํ”ˆ์†Œ์Šค ํ”„๋กœ์ ํŠธ์ด๋‹ค:

  • OPSET ์ •์˜: ์•ˆ์ •์ ์ธ ์—ฐ์‚ฐ ์„ธํŠธ๋ฅผ ์ •์˜ํ•˜์—ฌ ํ”„๋ ˆ์ž„์›Œํฌ ๊ฐ„ ์ด์‹์„ฑ ๋ณด์žฅ
  • Versioned Serialization: ๋ฒ„์ „ ๊ด€๋ฆฌ๊ฐ€ ๊ฐ€๋Šฅํ•œ ์ง๋ ฌํ™” ํ˜•์‹
  • ํฌํ„ฐ๋ธ”๋ฆฌํ‹ฐ: TensorFlow, JAX, PyTorch ๋“ฑ ๋‹ค์–‘ํ•œ ํ”„๋ ˆ์ž„์›Œํฌ์—์„œ ์‚ฌ์šฉ ๊ฐ€๋Šฅ
  • StableHLO Tooling: ๋ณ€ํ™˜, ๊ฒ€์ฆ, ์ง๋ ฌํ™” ๋„๊ตฌ ์ œ๊ณต

JIT ์ปดํŒŒ์ผ

XLA๋Š” JIT(Just-In-Time) ์ปดํŒŒ์ผ ๋ฐฉ์‹์„ ์ฑ„ํƒํ•˜์—ฌ ์‹คํ–‰ ์‹œ์ ์—์„œ ์ตœ์ ํ™”๋ฅผ ์ˆ˜ํ–‰ํ•œ๋‹ค:

  • ์ง€์—ฐ ์ปดํŒŒ์ผ(Lazy Compilation): ๋ชจ๋ธ ์‹คํ–‰ ์ง์ „์— ์ปดํŒŒ์ผ ์ˆ˜ํ–‰
  • ์บ์‹ฑ: ์ปดํŒŒ์ผ ๊ฒฐ๊ณผ๋ฅผ ์บ์‹ฑํ•˜์—ฌ ๋ฐ˜๋ณต ์‹คํ–‰ ์‹œ ์˜ค๋ฒ„ํ—ค๋“œ ์ตœ์†Œํ™”
  • ๋™์  shape ์ง€์›: ๋Ÿฐํƒ€์ž„์— ํฌ๊ธฐ๊ฐ€ ๊ฒฐ์ •๋˜๋Š” ํ…์„œ ์ฒ˜๋ฆฌ ๊ฐ€๋Šฅ

๋น„๊ต/๋ถ„์„

XLA vs ๋‹ค๋ฅธ AI ์ปดํŒŒ์ผ๋Ÿฌ

ํ•ญ๋ชฉ XLA TVM TensorRT Glow
๊ฐœ๋ฐœ์‚ฌ Google Apache NVIDIA Meta
๋ผ์ด์„ ์Šค Apache 2.0 Apache 2.0 ๋…์  BSD
์ง€์› ํ•˜๋“œ์›จ์–ด TPU, GPU, CPU ๋ฒ”์šฉ (GPU/CPU/NPU/TPU) NVIDIA GPU GPU, NPU
์ž๋™ ํŠœ๋‹ ์—†์Œ MetaSchedule ์—†์Œ ์—†์Œ
ํ”„๋ ˆ์ž„์›Œํฌ ์ง€์› TensorFlow, JAX ONNX, PyTorch, TF, MXNet TensorRT ํŒŒ์„œ PyTorch, Caffe2
LLM ์ง€์› StableHLO Relax + TIRx TensorRT-LLM ์ œํ•œ์ 
๋ถ„์‚ฐ ์‹คํ–‰ TPU Pod Disco (NCCL/RCCL) ๋‚ด์žฅ ์ œํ•œ์ 

XLA ์ตœ์ ํ™” ๊ธฐ๋ฒ• ๋น„๊ต

์ตœ์ ํ™” XLA TVM TensorRT
์—ฐ์‚ฐ ์œตํ•ฉ O O O
๋ฉ”๋ชจ๋ฆฌ ์ตœ์ ํ™” O O O
์ž๋™ ํŠœ๋‹ X O (MetaSchedule) X
๋™์  shape O O ์ œํ•œ์ 
์ •๋ฐ€๋„ ์ตœ์ ํ™” O O O (FP16/INT8)

ํ•˜๋“œ์›จ์–ด๋ณ„ XLA ์ง€์› ๋น„๊ต

ํ•˜๋“œ์›จ์–ด XLA ์ง€์› ํŠน์ง•
TPU ์™„์ „ ์ง€์› XLA์˜ ์ฃผ์š” ํƒ€๊ฒŸ, ์ตœ์ ํ™”๋œ PTX ์ƒ์„ฑ
GPU (NVIDIA) ์ง€์› CUDA/cuDNN ํ˜ธ์ถœ ์ฝ”๋“œ ์ƒ์„ฑ
GPU (AMD) ์ œํ•œ์  ROCm ๋ฐฑ์—”๋“œ ๊ฐœ๋ฐœ ์ค‘
CPU ์ง€์› LLVM ๊ธฐ๋ฐ˜ ๊ธฐ๊ณ„์–ด ์ƒ์„ฑ

๋™์ž‘ ์›๋ฆฌ

XLA ์ปดํŒŒ์ผ๋Ÿฌ ํŒŒ์ดํ”„๋ผ์ธ

XLA Compilation Pipeline
  1. ๋ชจ๋ธ ์ž„ํฌํŠธ ๋‹จ๊ณ„:
    - TensorFlow: tf.function(jit_compile=True) ๋ฐ์ฝ”๋ ˆ์ดํ„ฐ ์‚ฌ์šฉ
    - JAX: @jax.jit ๋ฐ์ฝ”๋ ˆ์ดํ„ฐ ์‚ฌ์šฉ
    - ํ”„๋ ˆ์ž„์›Œํฌ ๊ทธ๋ž˜ํ”„๋ฅผ HLO๋กœ ๋ณ€ํ™˜

  2. HLO ๊ทธ๋ž˜ํ”„ ์ตœ์ ํ™” ๋‹จ๊ณ„:
    - Constant Folding: ์ปดํŒŒ์ผ ์‹œ์ ์— ์ƒ์ˆ˜ ์—ฐ์‚ฐ ์ˆ˜ํ–‰
    - Dead Code Elimination: ๋ถˆํ•„์š”ํ•œ ์—ฐ์‚ฐ ์ œ๊ฑฐ
    - Operation Fusion: ๋…๋ฆฝ์ ์ธ ์—ฐ์‚ฐ ๊ฒฐํ•ฉ
    - Layout Optimization: ํ•˜๋“œ์›จ์–ด์— ์ตœ์ ๋œ ๋ฐ์ดํ„ฐ ๋ ˆ์ด์•„์›ƒ ๊ฒฐ์ •

  3. HLO ๋กœ์–ด๋ง ๋‹จ๊ณ„:
    - HLO ๋ช…๋ น์–ด๋ฅผ ํ•˜๋“œ์›จ์–ด๋ณ„ ์ค‘๊ฐ„ ํ‘œํ˜„์œผ๋กœ ๋ณ€ํ™˜
    - TPU: XLA-specific PTX ์ƒ์„ฑ
    - GPU: CUDA/cuDNN ํ˜ธ์ถœ ์ฝ”๋“œ ์ƒ์„ฑ
    - CPU: LLVM IR ์ƒ์„ฑ

  4. ์ฝ”๋“œ ์ƒ์„ฑ ๋‹จ๊ณ„:
    - ํ•˜๋“œ์›จ์–ด๋ณ„ ๊ธฐ๊ณ„์–ด ์ƒ์„ฑ
    - ๋ฉ”๋ชจ๋ฆฌ ํ• ๋‹น ๋ฐ ๊ด€๋ฆฌ ์ฝ”๋“œ ์ƒ์„ฑ
    - ๋Ÿฐํƒ€์ž„ ํ†ตํ•ฉ ์ฝ”๋“œ ์ƒ์„ฑ

  5. ๋Ÿฐํƒ€์ž„ ์‹คํ–‰ ๋‹จ๊ณ„:
    - ์ปดํŒŒ์ผ๋œ ์ปค๋„ ์‹คํ–‰
    - ๋ฉ”๋ชจ๋ฆฌ ๊ด€๋ฆฌ ๋ฐ ๋™๊ธฐํ™”
    - ๊ฒฐ๊ณผ ๋ฐ˜ํ™˜

HLO ์—ฐ์‚ฐ ์œตํ•ฉ

HLO Operation Fusion

HLO๋Š” ๋‹ค์–‘ํ•œ ์œตํ•ฉ ํŒจํ„ด์„ ์ง€์›ํ•œ๋‹ค:

์œตํ•ฉ ํŒจํ„ด ์„ค๋ช… ์˜ˆ์‹œ
์นดํ…Œ๊ณ ๋ฆฌ ์œตํ•ฉ ๋™์ผํ•œ ์นดํ…Œ๊ณ ๋ฆฌ์˜ ์—ฐ์‚ฐ ๊ฒฐํ•ฉ Conv + Bias + ReLU
์ˆ˜์ง ์œตํ•ฉ ๋ ˆ์ด์–ด ๊ฐ„ ์œตํ•ฉ Conv โ†’ Pool โ†’ BN
์ˆ˜ํ‰ ์œตํ•ฉ ๋™์ผ ์ž…๋ ฅ์„ ๊ณต์œ ํ•˜๋Š” ์—ฐ์‚ฐ ์œตํ•ฉ ๋ณ‘๋ ฌ ์ปจ๋ณผ๋ฃจ์…˜ ๋ธŒ๋žœ์น˜
์žฌ์‚ฌ์šฉ ์œตํ•ฉ ์ค‘๊ฐ„ ๊ฒฐ๊ณผ ์žฌ์‚ฌ์šฉ Attention ์Šค์ฝ”์–ด ๊ณ„์‚ฐ

TPU์™€์˜ ํ†ตํ•ฉ

XLA๋Š” TPU์™€ ๊ธด๋ฐ€ํ•˜๊ฒŒ ํ†ตํ•ฉ๋˜์–ด ์žˆ๋‹ค:

  • TPU Compiler: HLO๋ฅผ TPU-specific PTX๋กœ ๋ณ€ํ™˜
  • ๋ฉ”๋ชจ๋ฆฌ ๋ ˆ์ด์•„์›ƒ: TPU์˜ 2D ๋ฉ”๋ชจ๋ฆฌ ๊ตฌ์กฐ์— ์ตœ์ ํ™”๋œ ๋ ˆ์ด์•„์›ƒ ์ž๋™ ๊ฒฐ์ •
  • ์—ฐ์‚ฐ ์Šค์ผ€์ค„๋ง: TPU์˜ ๋งคํŠธ๋ฆญ์Šค ์œ ๋‹› ํ™œ์šฉ๋„๋ฅผ ๊ทน๋Œ€ํ™”ํ•˜๋Š” ์Šค์ผ€์ค„๋ง
  • ๋ถ„์‚ฐ ์‹คํ–‰: TPU Pod ๊ฐ„ ํ†ต์‹  ์ตœ์ ํ™”

์žฅ๋‹จ์ 

XLA ์žฅ๋‹จ์ 

์žฅ์  ๋‹จ์ 
TensorFlow/JAX ๊ธด๋ฐ€ ํ†ตํ•ฉ TensorFlow/JAX ์™ธ ์ง€์› ์ œํ•œ
TPU ์ตœ์ ํ™” GPU ์„ฑ๋Šฅ ์ œํ•œ
JIT ์ปดํŒŒ์ผ๋กœ ๋™์  shape ์ง€์› ๋Ÿฐํƒ€์ž„ ์ปดํŒŒ์ผ ์˜ค๋ฒ„ํ—ค๋“œ
StableHLO๋กœ ํ”„๋ ˆ์ž„์›Œํฌ ๋…๋ฆฝ์„ฑ ์ถ”๊ตฌ ์ž๋™ ํŠœ๋‹ ๊ธฐ๋Šฅ ๋ถ€์กฑ
Google์˜ ์ง€์†์ ์ธ ์ง€์› ์ปค๋ฎค๋‹ˆํ‹ฐ ์˜์กด
์•ˆ์ •์ ์ธ HLO OPSET ์ผ๋ถ€ ์ตœ์‹  ํ•˜๋“œ์›จ์–ด ์ง€์› ์ง€์—ฐ

XLA ์‚ฌ์šฉ ์‹œ ๊ณ ๋ ค์‚ฌํ•ญ

์ƒํ™ฉ ๊ถŒ์žฅ ์„ ํƒ
TensorFlow/JAX + TPU ์‚ฌ์šฉ ์‹œ XLA
NVIDIA GPU ์ถ”๋ก ๋งŒ ํ•„์š”ํ•œ ๊ฒฝ์šฐ TensorRT
๋‹ค์–‘ํ•œ ํ•˜๋“œ์›จ์–ด ํƒ€๊ฒŸ์ด ํ•„์š”ํ•œ ๊ฒฝ์šฐ TVM
์‹ค์‹œ๊ฐ„ ์ถ”๋ก  ์„ฑ๋Šฅ์ด ์ตœ์šฐ์„ ์ธ ๊ฒฝ์šฐ TensorRT ๋˜๋Š” ์ˆ˜๋™ ์ตœ์ ํ™”
์—ฐ๊ตฌ์šฉ ํ”„๋กœํ† ํƒ€์ดํ•‘ TVM (MetaSchedule ํ™œ์šฉ)
์—ฃ์ง€ ๋””๋ฐ”์ด์Šค ๋ฐฐํฌ TVM (ํฌ๋กœ์Šค ์ปดํŒŒ์ผ + RPC)

๊ด€๋ จ ๊ธฐ์ˆ 

๊ด€๋ จ ๋ฌธ์„œ

์ฐธ๊ณ  ๋ฌธํ—Œ

  • "XLA: Optimizing Compiler for Machine Learning," Google AI Blog, 2017
  • StableHLO Specification: https://github.com/openxla/stablehlo
  • "StableHLO: A Portable, High-Level ML Dialect," arXiv 2023
  • TensorFlow XLA Documentation: https://www.tensorflow.org/xla
  • JAX Documentation: https://jax.readthedocs.io/
  • "The MLIR Infrastructure: A Systematic Approach to ML Compiler Infrastructure," LLVM Developer Conference, 2022
  • Google TPU Documentation: https://cloud.google.com/tpu/docs

ํ•ต์‹ฌ ์ •๋ฆฌ

  1. XLA๋Š” Google์ด ๊ฐœ๋ฐœํ•œ ์˜คํ”ˆ์†Œ์Šค ML ์ปดํŒŒ์ผ๋Ÿฌ๋กœ, TensorFlow ๋ฐ JAX ๋ชจ๋ธ์„ TPU, GPU, CPU ๋“ฑ ๋‹ค์–‘ํ•œ ํ•˜๋“œ์›จ์–ด์—์„œ ํšจ์œจ์ ์œผ๋กœ ์‹คํ–‰ํ•˜๊ธฐ ์œ„ํ•ด ์ตœ์ ํ™”๋œ ์ฝ”๋“œ๋ฅผ ์ž๋™์œผ๋กœ ์ƒ์„ฑํ•œ๋‹ค
  2. HLO(High Level Optimizer) ์ค‘๊ฐ„ ํ‘œํ˜„์„ ํ†ตํ•ด ๊ทธ๋ž˜ํ”„ ์ˆ˜์ค€ ์ตœ์ ํ™”๋ฅผ ์ˆ˜ํ–‰ํ•˜๋ฉฐ, StableHLO๋ฅผ ํ†ตํ•ด ํ”„๋ ˆ์ž„์›Œํฌ ๋…๋ฆฝ์ ์ธ ํ˜ธํ™˜์„ฑ์„ ์ถ”๊ตฌํ•˜๊ณ  ์žˆ๋‹ค
  3. JIT ์ปดํŒŒ์ผ ๋ฐฉ์‹์„ ์ฑ„ํƒํ•˜์—ฌ ์‹คํ–‰ ์‹œ์ ์—์„œ ์ตœ์ ํ™”๋ฅผ ์ˆ˜ํ–‰ํ•˜๋ฉฐ, ์—ฐ์‚ฐ ์œตํ•ฉ, ๋ฉ”๋ชจ๋ฆฌ ์ตœ์ ํ™”, ๋ณ‘๋ ฌํ™” ๋“ฑ์„ ์ž๋™์œผ๋กœ ์ˆ˜ํ–‰ํ•œ๋‹ค
  4. TPU์™€ ํŠนํžˆ ๊ธด๋ฐ€ํ•˜๊ฒŒ ํ†ตํ•ฉ๋˜์–ด ์žˆ์œผ๋ฉฐ, TPU-specific PTX ์ƒ์„ฑ๊ณผ 2D ๋ฉ”๋ชจ๋ฆฌ ๋ ˆ์ด์•„์›ƒ ์ตœ์ ํ™”๋ฅผ ์ง€์›ํ•œ๋‹ค
  5. XLA์˜ ์ฃผ์š” ๊ฐ•์ ์€ TensorFlow/JAX ์ƒํƒœ๊ณ„์™€์˜ ํ†ตํ•ฉ์ด๋ฉฐ, ๋‹จ์ ์€ ํ”„๋ ˆ์ž„์›Œํฌ ์™ธ ์ง€์› ์ œํ•œ๊ณผ ์ž๋™ ํŠœ๋‹ ๊ธฐ๋Šฅ ๋ถ€์กฑ์ด๋‹ค