본문 바로가기
C.W.K.
Stream
Lesson 02 of 06 · published

Dtype: float32와 bfloat16, 정밀도의 함정

~10 min · numpy, jax, tutorial

Level 0호기심
0 XP0/73 lessons0/17 achievements
0/100 XP to next level100 XP to go0% complete

JAX의 기본 dtype은 float32이고 NumPy의 기본값은 float64야. 이 차이를 모르고 코드를 옮기면 같은 수식을 계산했는데도 결과가 달라질 수 있어.

import numpy as np
import jax.numpy as jnp

a_np = np.array([1.0, 2.0, 3.0])
a_jax = jnp.array([1.0, 2.0, 3.0])

print(a_np.dtype)   # float64
print(a_jax.dtype)  # float32  (!)

float32가 기본인 이유는 가속기의 특성과 관련이 있어. GPU와 TPU에서 float32 연산은 보통 float64보다 2~30배 빠르고, 대부분의 머신러닝 작업은 float32 정밀도로도 충분히 수렴해.

import jax
jax.config.update("jax_enable_x64", True)

import jax.numpy as jnp
a = jnp.array([1.0, 2.0])
print(a.dtype)  # float64 (이제 됨)

⚠️ 정밀도 함정

NumPy에서 float64로 계산하던 코드를 JAX로 옮기면 같은 결과가 나오지 않을 수 있어. 학습이 수렴하지 않거나 NaN이 나타나면 dtype부터 확인해.

Code

import jax.numpy as jnp

# Standard dtypes
f32 = jnp.array(1.0, dtype=jnp.float32)   # 32-bit float (default)
f64 = jnp.array(1.0, dtype=jnp.float64)   # 64-bit float
i32 = jnp.array(1, dtype=jnp.int32)       # 32-bit integer
b = jnp.array(True, dtype=jnp.bool_)      # Boolean

# ML-specific dtypes
f16 = jnp.array(1.0, dtype=jnp.float16)   # 16-bit float (half precision)
bf16 = jnp.array(1.0, dtype=jnp.bfloat16) # Brain floating point 16
import jax.numpy as jnp

# float16 has limited range
try:
    big_f16 = jnp.array(100000.0, dtype=jnp.float16)
    print(f"float16: {big_f16}")  # inf — overflows!
except:
    pass

# bfloat16 handles the same value fine
big_bf16 = jnp.array(100000.0, dtype=jnp.bfloat16)
print(f"bfloat16: {big_bf16}")  # 99840.0 — less precise but doesn't overflow

# float32 for reference
big_f32 = jnp.array(100000.0, dtype=jnp.float32)
print(f"float32: {big_f32}")    # 100000.0

# Memory usage: half of float32
arr_f32 = jnp.ones((1000, 1000), dtype=jnp.float32)
arr_bf16 = jnp.ones((1000, 1000), dtype=jnp.bfloat16)
print(f"float32 size: {arr_f32.nbytes / 1e6:.1f} MB")  # 4.0 MB
print(f"bfloat16 size: {arr_bf16.nbytes / 1e6:.1f} MB") # 2.0 MB
import jax
jax.config.update("jax_enable_x64", True)

# Now float64 is available
import jax.numpy as jnp
x = jnp.array(1.0, dtype=jnp.float64)
print(x.dtype)  # float64

External links

Exercise

같은 내적을 float32, x64를 활성화한 float64, bfloat16으로 계산해. 결과와 최대 절대 오차, 실행 시간을 출력하고 bfloat16에서 결과가 크게 무너지는 입력 하나를 찾아. 실제 학습 코드에서는 dtype을 어떻게 선택할지 적어.

Progress

Progress is local-only — sign in to sync across devices.
이 페이지에서 버그를 발견하셨거나 피드백이 있으세요?문제 신고

댓글 0

🔔 답글 알림 (로그인 필요)
로그인댓글을 남기려면 로그인해 주세요.

아직 댓글이 없어요. 첫 댓글을 남겨보세요.