Lisandro commited on
Commit
889f780
·
1 Parent(s): f331f9f

feat: Refactor optimize_pipeline_ function for improved dynamic shape handling and stability in ZeroGPU

Browse files
Files changed (1) hide show
  1. optimization.py +12 -11
optimization.py CHANGED
@@ -6,11 +6,9 @@ from torch.utils._pytree import tree_map
6
 
7
  P = ParamSpec("P")
8
 
9
- # Dimensiones fijas sugeridas por torch.export (evitan el UserError)
10
  TEXT_SEQ_LENGTH = 12
11
  IMAGE_SEQ_LENGTH = 4096
12
 
13
- # Configuraciones del compilador AOTInductor
14
  INDUCTOR_CONFIGS = {
15
  "conv_1x1_as_mm": True,
16
  "epilogue_fusion": False,
@@ -23,8 +21,10 @@ INDUCTOR_CONFIGS = {
23
 
24
  def optimize_pipeline_(pipeline: Callable[P, Any], *args: P.args, **kwargs: P.kwargs):
25
  """
26
- Compila el transformer de Qwen-Image usando AOTInductor para ZeroGPU.
27
- Versión adaptada específicamente a tu pipeline: sin dynamic_shapes conflictivos.
 
 
28
  """
29
 
30
  if not torch.cuda.is_available():
@@ -36,23 +36,24 @@ def optimize_pipeline_(pipeline: Callable[P, Any], *args: P.args, **kwargs: P.kw
36
  def compile_transformer():
37
  print("🏗️ Capturando modelo para AOT...")
38
 
39
- # Captura del grafo del transformer
40
  with spaces.aoti_capture(pipeline.transformer) as call:
41
  pipeline(*args, **kwargs)
42
 
43
- # Asignamos longitudes fijas a las secuencias (para evitar UserError)
44
  dynamic_shapes = tree_map(lambda t: None, call.kwargs)
45
- dynamic_shapes = {
46
- k: None for k in call.kwargs.keys()
47
- } # evitamos conflicto de claves
48
  static_shapes = {
49
  "hidden_states": {1: IMAGE_SEQ_LENGTH},
50
  "encoder_hidden_states": {1: TEXT_SEQ_LENGTH},
51
  "encoder_hidden_states_mask": {1: TEXT_SEQ_LENGTH},
 
 
52
  }
53
- # Solo agregamos las que existan
 
54
  for k, v in static_shapes.items():
55
- if k in dynamic_shapes:
56
  dynamic_shapes[k] = v
57
 
58
  print("🚀 Exportando modelo con torch.export...")
 
6
 
7
  P = ParamSpec("P")
8
 
 
9
  TEXT_SEQ_LENGTH = 12
10
  IMAGE_SEQ_LENGTH = 4096
11
 
 
12
  INDUCTOR_CONFIGS = {
13
  "conv_1x1_as_mm": True,
14
  "epilogue_fusion": False,
 
21
 
22
  def optimize_pipeline_(pipeline: Callable[P, Any], *args: P.args, **kwargs: P.kwargs):
23
  """
24
+ Versión estable para Qwen-Image:
25
+ - Corrige estructura de dynamic_shapes (incluye img_shapes como lista)
26
+ - Usa longitudes fijas sugeridas
27
+ - Funciona en ZeroGPU sin romper el Space
28
  """
29
 
30
  if not torch.cuda.is_available():
 
36
  def compile_transformer():
37
  print("🏗️ Capturando modelo para AOT...")
38
 
 
39
  with spaces.aoti_capture(pipeline.transformer) as call:
40
  pipeline(*args, **kwargs)
41
 
42
+ # Base dinámica (misma estructura que inputs)
43
  dynamic_shapes = tree_map(lambda t: None, call.kwargs)
44
+
45
+ # Ajustes estáticos seguros
 
46
  static_shapes = {
47
  "hidden_states": {1: IMAGE_SEQ_LENGTH},
48
  "encoder_hidden_states": {1: TEXT_SEQ_LENGTH},
49
  "encoder_hidden_states_mask": {1: TEXT_SEQ_LENGTH},
50
+ # img_shapes es una lista -> declaramos lista con None
51
+ "img_shapes": [None],
52
  }
53
+
54
+ # Aplicar solo los keys existentes
55
  for k, v in static_shapes.items():
56
+ if k in call.kwargs:
57
  dynamic_shapes[k] = v
58
 
59
  print("🚀 Exportando modelo con torch.export...")