텐서의 shape을 바꾸는 코드는 머신러닝 코드의 5~10%를 차지할 만큼 자주 등장해. 기본 연산에 익숙해질 필요가 있어. JAX의 reshape 방식은 NumPy와 같아.
import jax.numpy as jnp
a = jnp.arange(24)
b = a.reshape(2, 3, 4)
c = a.reshape(-1, 4) # -1 = "알아서 계산"
d = a.reshape(4, 6).T # transpose
e = jnp.expand_dims(a, axis=0)
f = jnp.squeeze(e)
import jax.numpy as jnp
# 2D transpose
a = jnp.array([[1, 2, 3], [4, 5, 6]])
print(a.T.shape) # (3, 2)
# Higher-dimensional: permute axes
# Common in ML: converting between channels-first and channels-last
img = jnp.zeros((32, 3, 224, 224)) # (batch, channels, height, width)
# NCHW -> NHWC
img_nhwc = jnp.transpose(img, (0, 2, 3, 1))
print(img_nhwc.shape) # (32, 224, 224, 3)
import jax.numpy as jnp
a = jnp.array([1, 2, 3])
b = jnp.array([4, 5, 6])
# Concatenate: join along existing axis
c = jnp.concatenate([a, b])
print(c) # [1 2 3 4 5 6]
# Stack: join along a NEW axis
s = jnp.stack([a, b])
print(s) # [[1 2 3], [4 5 6]]
print(s.shape) # (2, 3)
# vstack and hstack
v = jnp.vstack([a, b]) # Same as stack for 1D → 2D
h = jnp.hstack([a, b]) # Same as concatenate for 1D
(32, 28, 28, 3) 이미지 배치를 (32, 28*28*3)으로 평탄화한 뒤 원래 shape으로 복원해. 이어서 channels-first 형식으로 transpose하고 각 단계의 새 shape과 stride를 출력해. 이런 레이아웃 변환 연습은 신경망 코드의 조용한 shape 버그를 예방하는 가장 값싼 방법이야.
Progress
Progress is local-only — sign in to sync across devices.