Triton 언어 API — tl 연산 살펴보기

Triton 언어 API — tl 연산 살펴보기

Triton 커널 내부에서는 triton.language(약칭 tl)의 연산을 써요. 블록(텐서) 단위로 메모리 로드·스토어·수학 연산·리덕션·원자 연산을 다루는 기본 원시 연산들이 여기 모여 있어요.

출처: https://triton-lang.org/main/python-api/triton.language.html

프로그래밍 모델

  • tl.tensor — 값 또는 포인터의 N차원 배열
  • tl.tensor_descriptor — 전역 메모리 안의 텐서를 나타내는 서술자
  • tl.num_programs — 주어진 axis로 실행되는 프로그램 인스턴스 수

생성·모양 연산

  • tl.arange(start, end) — 반개구간 [start, end)의 연속 값
  • tl.full, tl.zeros_like, tl.cast, tl.to_tensor
  • tl.broadcast, tl.reshape, tl.trans, tl.permute, tl.split, tl.join

메모리/포인터 연산

  • tl.load(pointer, mask=...) — 포인터가 가리키는 메모리에서 데이터 로드
  • tl.store(pointer, value, mask=...) — 텐서를 메모리에 저장
  • tl.load_tensor_descriptor — 텐서 서술자에서 데이터 블록 로드

수학·리덕션 연산

  • tl.abs, tl.add, tl.sub, tl.cdiv, tl.log2, tl.exp, tl.sqrt
  • 리덕션: tl.sum, tl.min, tl.max, tl.argmax, tl.argmin, tl.reduce, tl.xor_sum

원자 연산

  • tl.atomic_add, tl.atomic_and, tl.atomic_xor — 지정 포인터 위치에서 원자적 연산 수행

난수 생성

  • tl.rand(seed, offset)U(0, 1)float32 블록
  • tl.randn(seed, offset)N(0, 1)float32 블록

더 알아보기