Trax 소개 — 깔끔한 코드와 속도의 딥러닝 라이브러리

Trax 소개

딥러닝 코드는 보통 너무 복잡해서 읽기 지치기 쉬워요. Trax는 "깔끔한 코드(clear code)"와 속도에 집중한 end-to-end 딥러닝 라이브러리예요. Google Brain 팀 내에서 적극적으로 사용·유지되고 있고, 기본 모델과 RL 알고리즘, 연구용 최신 모델까지 폭넓게 담고 있어요.

출처: https://github.com/google/trax

사전학습 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에서 그대로 동작해요.

더 알아보기