Spaces:
Running on Zero
Running on Zero
feat: Refactor optimize_pipeline_ function for improved dynamic shape handling and stability in ZeroGPU
Browse files- 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 |
-
|
| 27 |
-
|
|
|
|
|
|
|
| 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 |
-
#
|
| 44 |
dynamic_shapes = tree_map(lambda t: None, call.kwargs)
|
| 45 |
-
|
| 46 |
-
|
| 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 |
-
|
|
|
|
| 54 |
for k, v in static_shapes.items():
|
| 55 |
-
if k in
|
| 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...")
|