Lisandro commited on
Commit
da763f5
·
1 Parent(s): 58f8f91

feat: Refactor AOT optimization in pipeline to enhance dynamic shape handling and error management

Browse files
Files changed (1) hide show
  1. optimization.py +64 -44
optimization.py CHANGED
@@ -1,57 +1,77 @@
1
- from typing import Any
2
- from typing import Callable
3
- from typing import ParamSpec
4
- from torchao.quantization import quantize_
5
- from torchao.quantization import Float8DynamicActivationFloat8WeightConfig
6
  import spaces
7
  import torch
8
  from torch.utils._pytree import tree_map
9
 
10
- P = ParamSpec('P')
 
11
 
12
- TRANSFORMER_IMAGE_SEQ_LENGTH_DIM = torch.export.Dim('image_seq_length')
13
- TRANSFORMER_TEXT_SEQ_LENGTH_DIM = torch.export.Dim('text_seq_length')
 
14
 
 
15
  TRANSFORMER_DYNAMIC_SHAPES = {
16
- 'hidden_states': {
17
- 1: TRANSFORMER_IMAGE_SEQ_LENGTH_DIM,
18
- },
19
- 'encoder_hidden_states': {
20
- 1: TRANSFORMER_TEXT_SEQ_LENGTH_DIM,
21
- },
22
- 'encoder_hidden_states_mask': {
23
- 1: TRANSFORMER_TEXT_SEQ_LENGTH_DIM,
24
- },
25
- 'image_rotary_emb': ({
26
- 0: TRANSFORMER_IMAGE_SEQ_LENGTH_DIM,
27
- }, {
28
- 0: TRANSFORMER_TEXT_SEQ_LENGTH_DIM,
29
- }),
30
  }
31
 
 
32
  INDUCTOR_CONFIGS = {
33
- 'conv_1x1_as_mm': True,
34
- 'epilogue_fusion': False,
35
- 'coordinate_descent_tuning': True,
36
- 'coordinate_descent_check_all_directions': True,
37
- 'max_autotune': True,
38
- 'triton.cudagraphs': True,
39
  }
40
 
 
41
  def optimize_pipeline_(pipeline: Callable[P, Any], *args: P.args, **kwargs: P.kwargs):
42
- @spaces.GPU(duration=1500)
43
- def compile_transformer():
44
- with spaces.aoti_capture(pipeline.transformer) as call:
45
- pipeline(*args, **kwargs)
46
- dynamic_shapes = tree_map(lambda t: None, call.kwargs)
47
- dynamic_shapes |= TRANSFORMER_DYNAMIC_SHAPES
48
- # quantize_(pipeline.transformer, Float8DynamicActivationFloat8WeightConfig())
49
- exported = torch.export.export(
50
- mod=pipeline.transformer,
51
- args=call.args,
52
- kwargs=call.kwargs,
53
- dynamic_shapes=dynamic_shapes,
54
- )
55
- return spaces.aoti_compile(exported, INDUCTOR_CONFIGS)
56
-
57
- spaces.aoti_apply(compile_transformer(), pipeline.transformer)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # optimization.py
2
+ from typing import Any, Callable, ParamSpec
 
 
 
3
  import spaces
4
  import torch
5
  from torch.utils._pytree import tree_map
6
 
7
+ # No usamos torchao ni quantización, para evitar conflictos con versiones
8
+ P = ParamSpec("P")
9
 
10
+ # Definimos dimensiones dinámicas que sí existen en tu pipeline Qwen-Image
11
+ TRANSFORMER_IMAGE_SEQ_LENGTH_DIM = torch.export.Dim("image_seq_length")
12
+ TRANSFORMER_TEXT_SEQ_LENGTH_DIM = torch.export.Dim("text_seq_length")
13
 
14
+ # Solo incluimos las claves que realmente aparecen en tu modelo
15
  TRANSFORMER_DYNAMIC_SHAPES = {
16
+ "hidden_states": {1: TRANSFORMER_IMAGE_SEQ_LENGTH_DIM},
17
+ "encoder_hidden_states": {1: TRANSFORMER_TEXT_SEQ_LENGTH_DIM},
18
+ "encoder_hidden_states_mask": {1: TRANSFORMER_TEXT_SEQ_LENGTH_DIM},
 
 
 
 
 
 
 
 
 
 
 
19
  }
20
 
21
+ # Configuraciones del compilador AOTInductor
22
  INDUCTOR_CONFIGS = {
23
+ "conv_1x1_as_mm": True,
24
+ "epilogue_fusion": False,
25
+ "coordinate_descent_tuning": True,
26
+ "coordinate_descent_check_all_directions": True,
27
+ "max_autotune": True,
28
+ "triton.cudagraphs": True,
29
  }
30
 
31
+
32
  def optimize_pipeline_(pipeline: Callable[P, Any], *args: P.args, **kwargs: P.kwargs):
33
+ """
34
+ Optimiza el transformer interno del pipeline de Qwen-Image usando AOTInductor.
35
+ Funciona con ZeroGPU y solo compila el módulo transformer, no todo el pipeline.
36
+ """
37
+
38
+ # Si no hay GPU, no intentamos compilar
39
+ if not torch.cuda.is_available():
40
+ print("⚠️ CUDA no disponible. Se omite AOT.")
41
+ return pipeline
42
+
43
+ try:
44
+ @spaces.GPU(duration=1200)
45
+ def compile_transformer():
46
+ print("🏗️ Capturando modelo para AOT...")
47
+
48
+ # Capturamos la llamada del transformer
49
+ with spaces.aoti_capture(pipeline.transformer) as call:
50
+ pipeline(*args, **kwargs)
51
+
52
+ # Creamos el mapa de shapes dinámicos solo con claves válidas
53
+ dynamic_shapes = tree_map(lambda t: None, call.kwargs)
54
+ # Añadimos los shapes esperados, filtrando los que realmente existen
55
+ for k, v in TRANSFORMER_DYNAMIC_SHAPES.items():
56
+ if k in call.kwargs:
57
+ dynamic_shapes[k] = v
58
+
59
+ print("🚀 Exportando modelo con torch.export...")
60
+ exported = torch.export.export(
61
+ mod=pipeline.transformer,
62
+ args=call.args,
63
+ kwargs=call.kwargs,
64
+ dynamic_shapes=dynamic_shapes,
65
+ )
66
+
67
+ print("⚙️ Compilando con AOTInductor...")
68
+ return spaces.aoti_compile(exported, INDUCTOR_CONFIGS)
69
+
70
+ print("🧠 Aplicando AOT al transformer...")
71
+ spaces.aoti_apply(compile_transformer(), pipeline.transformer)
72
+ print("✅ AOT aplicado correctamente al transformer de Qwen-Image.")
73
+
74
+ except Exception as e:
75
+ print(f"⚠️ Error al aplicar AOT: {e}")
76
+
77
+ return pipeline