Torch FX로 모델 최적화하기 (변환 작성)

Torch FX로 모델 최적화하기 (변환 작성)

모델을 더 빠르게 만들려면 계산 그래프를 직접 바꿔야 할 때가 있어요. optimum.fx.optimization 모듈은 torch.fx 그래프 변환(transformation) 모음과, 여러분이 직접 변환을 작성하고 조합할 수 있는 클래스·함수를 제공해요. Optimum에는 되돌릴 수 있는 변환되돌릴 수 없는 변환 두 종류가 있어요. 하나씩 직접 만들어 보면서 감을 잡아볼게요.

출처: Hugging Face Optimum — Torch FX Optimization

되돌릴 수 없는 변환 작성하기

가장 기본적인 형태는 되돌릴 수 없는 변환이에요. 그래프 모듈에 적용하면 원래 모델을 되돌려받을 방법이 없는 변환이죠. 구현은 아주 쉬워요. [optimum.fx.optimization.Transformation]을 상속하고 transform() 메서드만 구현하면 돼요.

예를 들어 모든 곱셈을 덧셈으로 바꾸는 변환은 이렇게 생겼어요.

>>> import operator
>>> from optimum.fx.optimization import Transformation

>>> class ChangeMulToAdd(Transformation):
...     def transform(self, graph_module):
...         for node in graph_module.graph.nodes:
...             if node.op == "call_function" and node.target == operator.mul:
...                 node.target = operator.add
...         return graph_module

구현이 끝나면 이 변환을 일반 함수처럼 사용할 수 있어요.

>>> from transformers import BertModel
>>> from transformers.utils.fx import symbolic_trace

>>> model = BertModel.from_pretrained("bert-base-uncased")
>>> traced = symbolic_trace(
...     model,
...     input_names=["input_ids", "attention_mask", "token_type_ids"],
... )

>>> transformation = ChangeMulToAdd()
>>> transformed_model = transformation(traced)

되돌릴 수 있는 변환 작성하기

되돌릴 수 있는 변환은 변환과 그 역변환을 함께 구현해서, 변환된 모델에서 원래 모델을 되찾을 수 있게 해요. [optimum.fx.optimization.ReversibleTransformation]을 상속하고 transform()reverse()를 모두 구현해야 해요.

곱셈을 2배 곱셈으로 바꾸고, 역변환에서는 다시 2로 나누는 변환을 볼게요.

>>> import operator
>>> from optimum.fx.optimization import ReversibleTransformation

>>> class MulToMulTimesTwo(ReversibleTransformation):
...     def transform(self, graph_module):
...         for node in graph_module.graph.nodes:
...             if node.op == "call_function" and node.target == operator.mul:
...                 x, y = node.args
...                 node.args = (2 * x, y)
...         return graph_module
...
...     def reverse(self, graph_module):
...         for node in graph_module.graph.nodes:
...             if node.op == "call_function" and node.target == operator.mul:
...                 x, y = node.args
...                 node.args = (x / 2, y)
...         return graph_module

변환 조합하기

변환 하나만 쓰는 경우보다 여러 변환을 연쇄로 적용하는 일이 더 많아요. 이때 [optimum.fx.optimization.compose] 유틸리티 함수가 여러 변환을 이어붙여 하나의 변환으로 만들어줘요.

>>> from optimum.fx.optimization import compose
>>> composition = compose(MulToMulTimesTwo(), ChangeMulToAdd())

더 알아보기 (Learn more)