12 KiB
์ด๋ป๊ฒ ์ฌ์ฉ์ ์ ์ ํ์ดํ๋ผ์ธ์ ์์ฑํ๋์? how-to-create-a-custom-pipeline
์ด ๊ฐ์ด๋์์๋ ์ฌ์ฉ์ ์ ์ ํ์ดํ๋ผ์ธ์ ์ด๋ป๊ฒ ์์ฑํ๊ณ ํ๋ธ์ ๊ณต์ ํ๊ฑฐ๋ ๐ค Transformers ๋ผ์ด๋ธ๋ฌ๋ฆฌ์ ์ถ๊ฐํ๋ ๋ฐฉ๋ฒ์ ์ดํด๋ณด๊ฒ ์ต๋๋ค.
๋จผ์ ํ์ดํ๋ผ์ธ์ด ์์ฉํ ์ ์๋ ์์ ์
๋ ฅ์ ๊ฒฐ์ ํด์ผ ํฉ๋๋ค.
๋ฌธ์์ด, ์์ ๋ฐ์ดํธ, ๋์
๋๋ฆฌ ๋๋ ๊ฐ์ฅ ์ํ๋ ์
๋ ฅ์ผ ๊ฐ๋ฅ์ฑ์ด ๋์ ๊ฒ์ด๋ฉด ๋ฌด์์ด๋ ๊ฐ๋ฅํฉ๋๋ค.
์ด ์
๋ ฅ์ ๊ฐ๋ฅํ ํ ์์ํ Python ํ์์ผ๋ก ์ ์งํด์ผ (JSON์ ํตํด ๋ค๋ฅธ ์ธ์ด์๋) ํธํ์ฑ์ด ์ข์์ง๋๋ค.
์ด๊ฒ์ด ์ ์ฒ๋ฆฌ(preprocess) ํ์ดํ๋ผ์ธ์ ์
๋ ฅ(inputs)์ด ๋ ๊ฒ์
๋๋ค.
๊ทธ๋ฐ ๋ค์ outputs๋ฅผ ์ ์ํ์ธ์.
inputs์ ๊ฐ์ ์ ์ฑ
์ ๋ฐ๋ฅด๊ณ , ๊ฐ๋จํ ์๋ก ์ข์ต๋๋ค.
์ด๊ฒ์ด ํ์ฒ๋ฆฌ(postprocess) ๋ฉ์๋์ ์ถ๋ ฅ์ด ๋ ๊ฒ์
๋๋ค.
๋จผ์ 4๊ฐ์ ๋ฉ์๋(preprocess, _forward, postprocess ๋ฐ _sanitize_parameters)๋ฅผ ๊ตฌํํ๊ธฐ ์ํด ๊ธฐ๋ณธ ํด๋์ค Pipeline์ ์์ํ์ฌ ์์ํฉ๋๋ค.
from transformers import Pipeline
class MyPipeline(Pipeline):
def _sanitize_parameters(self, **kwargs):
preprocess_kwargs = {}
if "maybe_arg" in kwargs:
preprocess_kwargs["maybe_arg"] = kwargs["maybe_arg"]
return preprocess_kwargs, {}, {}
def preprocess(self, inputs, maybe_arg=2):
model_input = Tensor(inputs["input_ids"])
return {"model_input": model_input}
def _forward(self, model_inputs):
# model_inputs == {"model_input": model_input}
outputs = self.model(**model_inputs)
# Maybe {"logits": Tensor(...)}
return outputs
def postprocess(self, model_outputs):
best_class = model_outputs["logits"].softmax(-1)
return best_class
์ด ๋ถํ ๊ตฌ์กฐ๋ CPU/GPU์ ๋ํ ๋น๊ต์ ์ํํ ์ง์์ ์ ๊ณตํ๋ ๋์์, ๋ค๋ฅธ ์ค๋ ๋์์ CPU์ ๋ํ ์ฌ์ /์ฌํ ์ฒ๋ฆฌ๋ฅผ ์ํํ ์ ์๊ฒ ์ง์ํ๋ ๊ฒ์ ๋๋ค.
preprocess๋ ์๋ ์ ์๋ ์
๋ ฅ์ ๊ฐ์ ธ์ ๋ชจ๋ธ์ ๊ณต๊ธํ ์ ์๋ ํ์์ผ๋ก ๋ณํํฉ๋๋ค.
๋ ๋ง์ ์ ๋ณด๋ฅผ ํฌํจํ ์ ์์ผ๋ฉฐ ์ผ๋ฐ์ ์ผ๋ก Dict ํํ์
๋๋ค.
_forward๋ ๊ตฌํ ์ธ๋ถ ์ฌํญ์ด๋ฉฐ ์ง์ ํธ์ถํ ์ ์์ต๋๋ค.
forward๋ ์์ ์ฅ์น์์ ๋ชจ๋ ๊ฒ์ด ์๋ํ๋์ง ํ์ธํ๊ธฐ ์ํ ์์ ์ฅ์น๊ฐ ํฌํจ๋์ด ์์ด ์ ํธ๋๋ ํธ์ถ ๋ฉ์๋์
๋๋ค.
์ค์ ๋ชจ๋ธ๊ณผ ๊ด๋ จ๋ ๊ฒ์ _forward ๋ฉ์๋์ ์ํ๋ฉฐ, ๋๋จธ์ง๋ ์ ์ฒ๋ฆฌ/ํ์ฒ๋ฆฌ ๊ณผ์ ์ ์์ต๋๋ค.
postprocess ๋ฉ์๋๋ _forward์ ์ถ๋ ฅ์ ๊ฐ์ ธ์ ์ด์ ์ ๊ฒฐ์ ํ ์ต์ข
์ถ๋ ฅ ํ์์ผ๋ก ๋ณํํฉ๋๋ค.
_sanitize_parameters๋ ์ด๊ธฐํ ์๊ฐ์ pipeline(...., maybe_arg=4)์ด๋ ํธ์ถ ์๊ฐ์ pipe = pipeline(...); output = pipe(...., maybe_arg=4)๊ณผ ๊ฐ์ด, ์ฌ์ฉ์๊ฐ ์ํ๋ ๊ฒฝ์ฐ ์ธ์ ๋ ์ง ๋งค๊ฐ๋ณ์๋ฅผ ์ ๋ฌํ ์ ์๋๋ก ํ์ฉํฉ๋๋ค.
_sanitize_parameters์ ๋ฐํ ๊ฐ์ preprocess, _forward, postprocess์ ์ง์ ์ ๋ฌ๋๋ 3๊ฐ์ kwargs ๋์
๋๋ฆฌ์
๋๋ค.
ํธ์ถ์๊ฐ ์ถ๊ฐ ๋งค๊ฐ๋ณ์๋ก ํธ์ถํ์ง ์์๋ค๋ฉด ์๋ฌด๊ฒ๋ ์ฑ์ฐ์ง ๋ง์ญ์์ค.
์ด๋ ๊ฒ ํ๋ฉด ํญ์ ๋ "์์ฐ์ค๋ฌ์ด" ํจ์ ์ ์์ ๊ธฐ๋ณธ ์ธ์๋ฅผ ์ ์งํ ์ ์์ต๋๋ค.
๋ถ๋ฅ ์์
์์ top_k ๋งค๊ฐ๋ณ์๊ฐ ๋ํ์ ์ธ ์์
๋๋ค.
>>> pipe = pipeline("my-new-task")
>>> pipe("This is a test")
[{"label": "1-star", "score": 0.8}, {"label": "2-star", "score": 0.1}, {"label": "3-star", "score": 0.05}
{"label": "4-star", "score": 0.025}, {"label": "5-star", "score": 0.025}]
>>> pipe("This is a test", top_k=2)
[{"label": "1-star", "score": 0.8}, {"label": "2-star", "score": 0.1}]
์ด๋ฅผ ๋ฌ์ฑํ๊ธฐ ์ํด ์ฐ๋ฆฌ๋ postprocess ๋ฉ์๋๋ฅผ ๊ธฐ๋ณธ ๋งค๊ฐ๋ณ์์ธ 5๋ก ์
๋ฐ์ดํธํ๊ณ _sanitize_parameters๋ฅผ ์์ ํ์ฌ ์ด ์ ๋งค๊ฐ๋ณ์๋ฅผ ํ์ฉํฉ๋๋ค.
def postprocess(self, model_outputs, top_k=5):
best_class = model_outputs["logits"].softmax(-1)
# top_k๋ฅผ ์ฒ๋ฆฌํ๋ ๋ก์ง ์ถ๊ฐ
return best_class
def _sanitize_parameters(self, **kwargs):
preprocess_kwargs = {}
if "maybe_arg" in kwargs:
preprocess_kwargs["maybe_arg"] = kwargs["maybe_arg"]
postprocess_kwargs = {}
if "top_k" in kwargs:
postprocess_kwargs["top_k"] = kwargs["top_k"]
return preprocess_kwargs, {}, postprocess_kwargs
์ /์ถ๋ ฅ์ ๊ฐ๋ฅํ ํ ๊ฐ๋จํ๊ณ ์์ ํ JSON ์ง๋ ฌํ ๊ฐ๋ฅํ ํ์์ผ๋ก ์ ์งํ๋ ค๊ณ ๋ ธ๋ ฅํ์ญ์์ค. ์ด๋ ๊ฒ ํ๋ฉด ์ฌ์ฉ์๊ฐ ์๋ก์ด ์ข ๋ฅ์ ๊ฐ์ฒด๋ฅผ ์ดํดํ์ง ์๊ณ ๋ ํ์ดํ๋ผ์ธ์ ์ฝ๊ฒ ์ฌ์ฉํ ์ ์์ต๋๋ค. ๋ํ ์ฌ์ฉ ์ฉ์ด์ฑ์ ์ํด ์ฌ๋ฌ ๊ฐ์ง ์ ํ์ ์ธ์(์ค๋์ค ํ์ผ์ ํ์ผ ์ด๋ฆ, URL ๋๋ ์์ํ ๋ฐ์ดํธ์ผ ์ ์์)๋ฅผ ์ง์ํ๋ ๊ฒ์ด ๋น๊ต์ ์ผ๋ฐ์ ์ ๋๋ค.
์ง์๋๋ ์์ ๋ชฉ๋ก์ ์ถ๊ฐํ๊ธฐ adding-it-to-the-list-of-supported-tasks
new-task๋ฅผ ์ง์๋๋ ์์
๋ชฉ๋ก์ ๋ฑ๋กํ๋ ค๋ฉด PIPELINE_REGISTRY์ ์ถ๊ฐํด์ผ ํฉ๋๋ค:
from transformers.pipelines import PIPELINE_REGISTRY
PIPELINE_REGISTRY.register_pipeline(
"new-task",
pipeline_class=MyPipeline,
pt_model=AutoModelForSequenceClassification,
)
์ํ๋ ๊ฒฝ์ฐ ๊ธฐ๋ณธ ๋ชจ๋ธ์ ์ง์ ํ ์ ์์ผ๋ฉฐ, ์ด ๊ฒฝ์ฐ ํน์ ๊ฐ์ (๋ถ๊ธฐ ์ด๋ฆ ๋๋ ์ปค๋ฐ ํด์์ผ ์ ์์, ์ฌ๊ธฐ์๋ "abcdef")๊ณผ ํ์ ์ ํจ๊ป ๊ฐ์ ธ์์ผ ํฉ๋๋ค:
PIPELINE_REGISTRY.register_pipeline(
"new-task",
pipeline_class=MyPipeline,
pt_model=AutoModelForSequenceClassification,
default={"pt": ("user/awesome_model", "abcdef")},
type="text", # ํ์ฌ ์ง์ ์ ํ: text, audio, image, multimodal
)
Hub์ ํ์ดํ๋ผ์ธ ๊ณต์ ํ๊ธฐ share-your-pipeline-on-the-hub
Hub์ ์ฌ์ฉ์ ์ ์ ํ์ดํ๋ผ์ธ์ ๊ณต์ ํ๋ ค๋ฉด Pipeline ํ์ ํด๋์ค์ ์ฌ์ฉ์ ์ ์ ์ฝ๋๋ฅผ Python ํ์ผ์ ์ ์ฅํ๊ธฐ๋ง ํ๋ฉด ๋ฉ๋๋ค.
์๋ฅผ ๋ค์ด, ๋ค์๊ณผ ๊ฐ์ด ๋ฌธ์ฅ ์ ๋ถ๋ฅ๋ฅผ ์ํ ์ฌ์ฉ์ ์ ์ ํ์ดํ๋ผ์ธ์ ์ฌ์ฉํ๋ค๊ณ ๊ฐ์ ํด ๋ณด๊ฒ ์ต๋๋ค:
import numpy as np
from transformers import Pipeline
def softmax(outputs):
maxes = np.max(outputs, axis=-1, keepdims=True)
shifted_exp = np.exp(outputs - maxes)
return shifted_exp / shifted_exp.sum(axis=-1, keepdims=True)
class PairClassificationPipeline(Pipeline):
def _sanitize_parameters(self, **kwargs):
preprocess_kwargs = {}
if "second_text" in kwargs:
preprocess_kwargs["second_text"] = kwargs["second_text"]
return preprocess_kwargs, {}, {}
def preprocess(self, text, second_text=None):
return self.tokenizer(text, text_pair=second_text, return_tensors=self.framework)
def _forward(self, model_inputs):
return self.model(**model_inputs)
def postprocess(self, model_outputs):
logits = model_outputs.logits[0].numpy()
probabilities = softmax(logits)
best_class = np.argmax(probabilities)
label = self.model.config.id2label[best_class]
score = probabilities[best_class].item()
logits = logits.tolist()
return {"label": label, "score": score, "logits": logits}
๊ตฌํ์ ํ๋ ์์ํฌ์ ๊ตฌ์ ๋ฐ์ง ์์ผ๋ฉฐ, PyTorch์ TensorFlow ๋ชจ๋ธ์ ๋ํด ์๋ํฉ๋๋ค.
์ด๋ฅผ pair_classification.py๋ผ๋ ํ์ผ์ ์ ์ฅํ ๊ฒฝ์ฐ, ๋ค์๊ณผ ๊ฐ์ด ๊ฐ์ ธ์ค๊ณ ๋ฑ๋กํ ์ ์์ต๋๋ค:
from pair_classification import PairClassificationPipeline
from transformers.pipelines import PIPELINE_REGISTRY
from transformers import AutoModelForSequenceClassification, TFAutoModelForSequenceClassification
PIPELINE_REGISTRY.register_pipeline(
"pair-classification",
pipeline_class=PairClassificationPipeline,
pt_model=AutoModelForSequenceClassification,
tf_model=TFAutoModelForSequenceClassification,
)
์ด ์์
์ด ์๋ฃ๋๋ฉด ์ฌ์ ํ๋ จ๋ ๋ชจ๋ธ๊ณผ ํจ๊ป ์ฌ์ฉํ ์ ์์ต๋๋ค.
์๋ฅผ ๋ค์ด, sgugger/finetuned-bert-mrpc์ MRPC ๋ฐ์ดํฐ ์ธํธ์์ ๋ฏธ์ธ ์กฐ์ ๋์ด ๋ฌธ์ฅ ์์ ํจ๋ฌํ๋ ์ด์ฆ์ธ์ง ์๋์ง๋ฅผ ๋ถ๋ฅํฉ๋๋ค.
from transformers import pipeline
classifier = pipeline("pair-classification", model="sgugger/finetuned-bert-mrpc")
๊ทธ๋ฐ ๋ค์ push_to_hub ๋ฉ์๋๋ฅผ ์ฌ์ฉํ์ฌ ํ๋ธ์ ๊ณต์ ํ ์ ์์ต๋๋ค:
classifier.push_to_hub("test-dynamic-pipeline")
์ด๋ ๊ฒ ํ๋ฉด "test-dynamic-pipeline" ํด๋ ๋ด์ PairClassificationPipeline์ ์ ์ํ ํ์ผ์ด ๋ณต์ฌ๋๋ฉฐ, ํ์ดํ๋ผ์ธ์ ๋ชจ๋ธ๊ณผ ํ ํฌ๋์ด์ ๋ ์ ์ฅํ ํ, {your_username}/test-dynamic-pipeline ์ ์ฅ์์ ์๋ ๋ชจ๋ ๊ฒ์ ํธ์ํฉ๋๋ค.
์ดํ์๋ trust_remote_code=True ์ต์
๋ง ์ ๊ณตํ๋ฉด ๋๊ตฌ๋ ์ฌ์ฉํ ์ ์์ต๋๋ค.
from transformers import pipeline
classifier = pipeline(model="{your_username}/test-dynamic-pipeline", trust_remote_code=True)
๐ค Transformers์ ํ์ดํ๋ผ์ธ ์ถ๊ฐํ๊ธฐ add-the-pipeline-to-transformers
๐ค Transformers์ ์ฌ์ฉ์ ์ ์ ํ์ดํ๋ผ์ธ์ ๊ธฐ์ฌํ๋ ค๋ฉด, pipelines ํ์ ๋ชจ๋์ ์ฌ์ฉ์ ์ ์ ํ์ดํ๋ผ์ธ ์ฝ๋์ ํจ๊ป ์ ๋ชจ๋์ ์ถ๊ฐํ ๋ค์, pipelines/__init__.py์์ ์ ์๋ ์์
๋ชฉ๋ก์ ์ถ๊ฐํด์ผ ํฉ๋๋ค.
๊ทธ๋ฐ ๋ค์ ํ
์คํธ๋ฅผ ์ถ๊ฐํด์ผ ํฉ๋๋ค.
tests/test_pipelines_MY_PIPELINE.py๋ผ๋ ์ ํ์ผ์ ๋ง๋ค๊ณ ๋ค๋ฅธ ํ
์คํธ์ ์์ ๋ฅผ ํจ๊ป ์์ฑํฉ๋๋ค.
run_pipeline_test ํจ์๋ ๋งค์ฐ ์ผ๋ฐ์ ์ด๋ฉฐ, model_mapping ๋ฐ tf_model_mapping์์ ์ ์๋ ๊ฐ๋ฅํ ๋ชจ๋ ์ํคํ
์ฒ์ ์์ ๋ฌด์์ ๋ชจ๋ธ์์ ์คํ๋ฉ๋๋ค.
์ด๋ ํฅํ ํธํ์ฑ์ ํ
์คํธํ๋ ๋ฐ ๋งค์ฐ ์ค์ํ๋ฉฐ, ๋๊ตฐ๊ฐ XXXForQuestionAnswering์ ์ํ ์ ๋ชจ๋ธ์ ์ถ๊ฐํ๋ฉด ํ์ดํ๋ผ์ธ ํ
์คํธ๊ฐ ํด๋น ๋ชจ๋ธ์์ ์คํ์ ์๋ํ๋ค๋ ์๋ฏธ์
๋๋ค.
๋ชจ๋ธ์ด ๋ฌด์์์ด๊ธฐ ๋๋ฌธ์ ์ค์ ๊ฐ์ ํ์ธํ๋ ๊ฒ์ ๋ถ๊ฐ๋ฅํ๋ฏ๋ก, ๋จ์ํ ํ์ดํ๋ผ์ธ ์ถ๋ ฅ TYPE๊ณผ ์ผ์น์ํค๊ธฐ ์ํ ๋์ฐ๋ฏธ ANY๊ฐ ์์ต๋๋ค.
๋ํ 2๊ฐ(์ด์์ ์ผ๋ก๋ 4๊ฐ)์ ํ ์คํธ๋ฅผ ๊ตฌํํด์ผ ํฉ๋๋ค.
test_small_model_pt: ์ด ํ์ดํ๋ผ์ธ์ ๋ํ ์์ ๋ชจ๋ธ 1๊ฐ๋ฅผ ์ ์(๊ฒฐ๊ณผ๊ฐ ์๋ฏธ ์์ด๋ ์๊ด์์)ํ๊ณ ํ์ดํ๋ผ์ธ ์ถ๋ ฅ์ ํ ์คํธํฉ๋๋ค. ๊ฒฐ๊ณผ๋test_small_model_tf์ ๋์ผํด์ผ ํฉ๋๋ค.test_small_model_tf: ์ด ํ์ดํ๋ผ์ธ์ ๋ํ ์์ ๋ชจ๋ธ 1๊ฐ๋ฅผ ์ ์(๊ฒฐ๊ณผ๊ฐ ์๋ฏธ ์์ด๋ ์๊ด์์)ํ๊ณ ํ์ดํ๋ผ์ธ ์ถ๋ ฅ์ ํ ์คํธํฉ๋๋ค. ๊ฒฐ๊ณผ๋test_small_model_pt์ ๋์ผํด์ผ ํฉ๋๋ค.test_large_model_pt(์ ํ์ฌํญ): ๊ฒฐ๊ณผ๊ฐ ์๋ฏธ ์์ ๊ฒ์ผ๋ก ์์๋๋ ์ค์ ํ์ดํ๋ผ์ธ์์ ํ์ดํ๋ผ์ธ์ ํ ์คํธํฉ๋๋ค. ์ด๋ฌํ ํ ์คํธ๋ ์๋๊ฐ ๋๋ฆฌ๋ฏ๋ก ์ด๋ฅผ ํ์ํด์ผ ํฉ๋๋ค. ์ฌ๊ธฐ์์ ๋ชฉํ๋ ํ์ดํ๋ผ์ธ์ ๋ณด์ฌ์ฃผ๊ณ ํฅํ ๋ฆด๋ฆฌ์ฆ์์์ ๋ณํ๊ฐ ์๋์ง ํ์ธํ๋ ๊ฒ์ ๋๋ค.test_large_model_tf(์ ํ์ฌํญ): ๊ฒฐ๊ณผ๊ฐ ์๋ฏธ ์์ ๊ฒ์ผ๋ก ์์๋๋ ์ค์ ํ์ดํ๋ผ์ธ์์ ํ์ดํ๋ผ์ธ์ ํ ์คํธํฉ๋๋ค. ์ด๋ฌํ ํ ์คํธ๋ ์๋๊ฐ ๋๋ฆฌ๋ฏ๋ก ์ด๋ฅผ ํ์ํด์ผ ํฉ๋๋ค. ์ฌ๊ธฐ์์ ๋ชฉํ๋ ํ์ดํ๋ผ์ธ์ ๋ณด์ฌ์ฃผ๊ณ ํฅํ ๋ฆด๋ฆฌ์ฆ์์์ ๋ณํ๊ฐ ์๋์ง ํ์ธํ๋ ๊ฒ์ ๋๋ค.