ktrain 시작하기 (BERT 분류)
ktrain 시작하기 (BERT 분류)
ktrain으로 텍스트 분류를 시작하는 흐름은 반복적이에요. 데이터를 준비하고, Transformer로 전처리하고, 학습자(Learner)로 감싸서 fit하면 돼요. BERT 기반 분류도 몇 단계면 충분하죠.
설치부터 볼게요.
pip3 install ktrain
감성 분석(IMDb) 같은 텍스트 분류 예시를 보면 아래와 같은 구조예요. 먼저 트레이닝/테스트 데이터를 준비하고, ktrain의 text 모듈로 Transformer를 만드세요.
import ktrain
from ktrain import text
MODEL_NAME = 'distilbert-base-uncased'
t = text.Transformer(MODEL_NAME, maxlen=500, class_names=train_b.target_names)
trn = t.preprocess_train(part1_x, part1_y)
model = t.get_classifier()
learner = ktrain.get_learner(model, train_data=trn, val_data=None, batch_size=6)
이렇게 만들어진 learner가 학습의 중심이에요. learner.fit_onecycle(2e-5, 1)처럼 1cycle 정책으로 학습하거나, learner.fit(1e-3, 1)처럼 일반 방식으로 학습할 수 있어요.
예측과 배포는 get_predictor로 간단히 끝나요.
predictor = ktrain.get_predictor(learner.model, preproc)
이렇게 저장한 predictor는 predictor.predict(text)처럼 새 원본 데이터에 바로 적용할 수 있어요. ktrain은 모델 자체를 저장하는 데 그치지 않고 전처리 단계까지 함께 저장해서, 배포 시점에 원시 데이터를 그대로 넣어도 동작하게 해 주는 게 특징이에요.