Torch FX로 모델 최적화하기 (변환 작성)
Torch FX로 모델 최적화하기 (변환 작성)
모델을 더 빠르게 만들려면 계산 그래프를 직접 바꿔야 할 때가 있어요. optimum.fx.optimization 모듈은 torch.fx 그래프 변환(transformation) 모음과, 여러분이 직접 변환을 작성하고 조합할 수 있는 클래스·함수를 제공해요. Optimum에는 되돌릴 수 있는 변환과 되돌릴 수 없는 변환 두 종류가 있어요. 하나씩 직접 만들어 보면서 감을 잡아볼게요.
되돌릴 수 없는 변환 작성하기
가장 기본적인 형태는 되돌릴 수 없는 변환이에요. 그래프 모듈에 적용하면 원래 모델을 되돌려받을 방법이 없는 변환이죠. 구현은 아주 쉬워요. [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())