Terrantula/Chibi_Coloring_Book
0
1"""2"""3 4from typing import Any5from typing import Callable6from typing import ParamSpec7from torchao.quantization import quantize_8from torchao.quantization import Float8DynamicActivationFloat8WeightConfig9import spaces10import torch11from torch.utils._pytree import tree_map12 13 14P = ParamSpec('P')15 16 17TRANSFORMER_IMAGE_SEQ_LENGTH_DIM = torch.export.Dim('image_seq_length')18TRANSFORMER_TEXT_SEQ_LENGTH_DIM = torch.export.Dim('text_seq_length')19 20TRANSFORMER_DYNAMIC_SHAPES = {21 'hidden_states': {22 1: TRANSFORMER_IMAGE_SEQ_LENGTH_DIM,23 },24 'encoder_hidden_states': {25 1: TRANSFORMER_TEXT_SEQ_LENGTH_DIM,26 },27 'encoder_hidden_states_mask': {28 1: TRANSFORMER_TEXT_SEQ_LENGTH_DIM,29 },30 'image_rotary_emb': ({31 0: TRANSFORMER_IMAGE_SEQ_LENGTH_DIM,32 }, {33 0: TRANSFORMER_TEXT_SEQ_LENGTH_DIM,34 }),35}36 37 38INDUCTOR_CONFIGS = {39 'conv_1x1_as_mm': True,40 'epilogue_fusion': False,41 'coordinate_descent_tuning': True,42 'coordinate_descent_check_all_directions': True,43 'max_autotune': True,44 'triton.cudagraphs': True,45}46 47 48def optimize_pipeline_(pipeline: Callable[P, Any], *args: P.args, **kwargs: P.kwargs):49 50 @spaces.GPU(duration=1500)51 def compile_transformer():52 53 with spaces.aoti_capture(pipeline.transformer) as call:54 pipeline(*args, **kwargs)55 56 dynamic_shapes = tree_map(lambda t: None, call.kwargs)57 dynamic_shapes |= TRANSFORMER_DYNAMIC_SHAPES58 59 # quantize_(pipeline.transformer, Float8DynamicActivationFloat8WeightConfig())60 61 exported = torch.export.export(62 mod=pipeline.transformer,63 args=call.args,64 kwargs=call.kwargs,65 dynamic_shapes=dynamic_shapes,66 )67 68 return spaces.aoti_compile(exported, INDUCTOR_CONFIGS)69 70 spaces.aoti_apply(compile_transformer(), pipeline.transformer)