Trax 소개 — 깔끔한 코드와 속도의 딥러닝 라이브러리
Trax 소개
딥러닝 코드는 보통 너무 복잡해서 읽기 지치기 쉬워요. Trax는 "깔끔한 코드(clear code)"와 속도에 집중한 end-to-end 딥러닝 라이브러리예요. Google Brain 팀 내에서 적극적으로 사용·유지되고 있고, 기본 모델과 RL 알고리즘, 연구용 최신 모델까지 폭넓게 담고 있어요.
사전학습 Transformer로 번역기 만들기
몇 줄로 영어-독일어 번역기를 만들 수 있어요. pre-trained 가중치는 gs://trax-ml/models/translation/ende_wmt32k.gin 설정을 써요.
import trax
# Pre-trained 모델 설정: gs://trax-ml/models/translation/ende_wmt32k.gin
model = trax.models.Transformer(
input_vocab_size=33300, d_model=512, d_ff=2048, n_heads=8,
n_encoder_layers=6, n_decoder_layers=6, max_len=2048, mode='predict')
# 사전학습 가중치로 초기화
model.init_from_file('gs://trax-ml/models/translation/ende_wmt32k.pkl.gz', weights_only=True)
# 문장 토큰화
sentence = 'It is nice to learn new things today!'
tokenized = list(trax.data.tokenize(
iter([sentence]), # 스트림 위에서 동작
vocab_dir='gs://trax-ml/vocabs/', vocab_file='ende_32k.subword'))[0]
# Transformer로 디코딩
tokenized = tokenized[None, :] # 배치 차원 추가
tokenized_translation = trax.supervised.decoding.autoregressive_sample(
model, tokenized, temperature=0.0) # 높은 temperature: 더 다양한 결과
# 디토큰화
tokenized_translation = tokenized_translation[0][:-1] # 배치·EOS 제거
translation = trax.data.detokenize(tokenized_translation,
vocab_dir='gs://trax-ml/vocabs/', vocab_file='ende_32k.subword')
print(translation)
Es ist schön, heute neue Dinge zu lernen!
포함된 것
- 기본 모델: ResNet, LSTM, Transformer
- RL 알고리즘: REINFORCE, A2C, PPO
- 연구용 새 모델: Reformer
- 새 RL 알고리즘: AWR
- Tensor2Tensor, TensorFlow Datasets 등 다수의 데이터셋 바인딩
라이브러리로도, 셸 바이너리로도 쓸 수 있고 CPU·GPU·TPU에서 그대로 동작해요.