/
https://42jerrykim.github.io/ _index.md
NumPy 호환 배열에 자동미분과 XLA 컴파일을 결합한 JAX를 다루는 시리즈의 도입 챕터입니다. JAX가 존재하는 세 가지 이유부터 jit·grad·vmap·pmap을 아우르는 8개 챕터 커리큘럼, 학습 목표, PyTorch와의 관계까지 자세히 정리합니다. jax.numpy 배열이 NumPy와 API는 같지만 불변인 이유와 .at[idx].set() 사용법, 그리고 jit·grad 대상 함수가 순수해야 하는 이유를 전역 변수·리스트 append 오작동 예제와 lax.scan 수정 코드로 정리합니다. jax.jit이 함수를 트레이싱해 jaxpr이라는 중간 표현을 만들고 이를 XLA 컴파일러로 넘겨 기계어로 바꾸는 과정과, 재컴파일이 일어나는 조건·static_argnums·if/for 제어흐름 함정을 코드와 함께 상세히 정리합니다. jax.grad가 코드를 재실행하지 않고 정확한 도함수를 얻는 원리를 수치미분·기호미분과 대조해 설명합니다. forward/reverse-mode 차이, jacfwd·jacrev·custom_jvp를 손으로 검산한 코드로 다룹니다. vmap이 for 반복문 없이 배치 처리를 만드는 배칭 규칙과 in_axes·out_axes 지정법, vmap(grad(f))로 샘플별 기울기를 구하는 조합 패턴, pmap의 현재 상태와 권장 대안(jit·shard_map), pytree 개념을 정리합니다. jax.random.key와 split이 numpy.random 같은 전역 난수 상태 대신 필요한 이유를 키 재사용 버그로 재현하고, 이 상태 전달 방식이 04장 pytree·Optax opt_state와 동일한 원칙임을 다룹니다. Flax NNX가 2026년 현재 공식 권장 API임을 문서로 확인하고 nnx.Module로 MLP를 정의한 뒤, nnx.split과 jax.grad, Optax adam을 결합해 opt_state를 명시적으로 주고받는 jit 학습 스텝을 만드는 법을 다룹니다. torch.Tensor·autograd·nn.Module·optimizer.step()을 JAX 개념으로 매핑하고, 순수성 위반·전역 PRNG·jit 재컴파일 함정 3가지와 eager 대 trace-and-compile 실행 모델 차이를 실전 코드로 다룹니다. Pragmatic Engineer의 2026년 8월 인터뷰를 계기로, Casey Muratori가 2023년 실측 벤치마크로 보여준 클린 코드와 성능의 관계를 정리한다. 가상 함수 기반 다형성을 switch문·룩업 테이블로 바꾸면 왜 최대 25배 빨라지는지, Uncle Bob의 반박까지 함께 다룬다. PowerShell 121챕터 커리큘럼의 과정 개요. 텍스트가 아닌 객체 파이프라인이라는 정신 모델, 18개 Part 학습 순서의 설계 근거, 00-120장 전체 목차, 선수 지식과 완주 후 실무자가 갖추는 원격 관리·보안 역량까지 정리한 과정 개요 챕터다.