PyTorch-FX Shaper
El servidor MCP proporciona la forma de los tensores para convertir código PyTorch a einsum y einops.
Documentación
Agent Shaper
Agent Shaper extrae metadatos de forma de tensores por módulo de cualquier PyTorch nn.Module y los utiliza para anotar archivos fuente — ya sea con comentarios descriptivos de forma o reescribiendo operaciones como torch.einsum / einops.
Cómo funciona
- Extracción de formas (
fx_utils/get_fx_data.py) — ejecutatorch.export+ShapePropen tu módulo para capturar la forma de cada tensor intermedio en el forward pass de cada módulo definido en el workspace, no solo las entradas. - Anotación manual (
fx_utils/manual_annotate.py) — inserta las formas extraídas como comentarios en línea al final de cada línea de código relevante. - Anotación LLM (
fx_utils/llm_annotate.py) — alimenta cada clase de módulo anotada manualmente a un LLM en paralelo. Dos modos:- COMMENT — reescribe comentarios con nombres descriptivos de dimensiones (p. ej.
batch_size,seq_len,n_embd) y explicaciones en lenguaje natural de cada transformación. - EINSUM — reescribe el módulo completo reemplazando matmuls y operaciones de atención con
torch.einsum, colapsando reshapes intermedios cuando sea posible.
- COMMENT — reescribe comentarios con nombres descriptivos de dimensiones (p. ej.
- Visor de diffs — revisa los cambios del LLM antes de aceptarlos, ya sea en una interfaz Streamlit o directamente en el editor de diff nativo de VS Code.
- Servidor MCP (
agent_shaper/mcp_server.py) — expone el pipeline completo de extracción de formas y validación de reescrituras como herramientas MCP consumibles por Claude Code o cualquier asistente de IA compatible con MCP.
Todos los módulos en un archivo se procesan en paralelo mediante asyncio. Todo lo que está fuera de las clases de módulos (imports, dataclasses, objetos de configuración) se conserva sin cambios en el archivo de salida.
Estructura del proyecto
agent_shaper/
fx_utils/
get_fx_data.py # shape extraction via torch.export + ShapeProp
manual_annotate.py # inline shape comment insertion
llm_annotate.py # LLM-powered rewrite (COMMENT or EINSUM mode)
diff_viewer.py # Streamlit diff UI
mcp_server.py # MCP tool server for AI-assisted rewrites
examples/
transformer/ # GPT-2, LLaMA, Qwen3, Swin Transformer
model.py # nanoGPT reference implementation
model_einsum.py # einsum-rewritten version
llama.py / llama_einsum.py
qwen3.py / qwen3_einsum.py
swin_transformer.py / swin_transformer_einsum.py
run_transformer.py # example forward pass
alignment/ # standalone alignment loss references
alignment_losses.py # DPO/IPO/SimPO losses (reference einsum style)
dpo_losses_einsum.py
jsd.py / jsd_einsum.py
opd.py / opd_einsum.py
direct_alignment/ # full training-ready RLHF loss implementations
loss.py # DPO, IPO, SimPO, ORPO, KTO, APO-zero, APO-down
train.py # training loop
data.py # preference dataset loading
config.py # training configuration
Cada archivo *_einsum.py es la contraparte reescrita con einsum/einops del original, validada para producir salidas numéricamente idénticas (atol=1e-5).
Instalación
Requiere Python 3.12+.
python -m venv .venv
source .venv/bin/activate
pip install -e .
Uso
Servidor MCP (recomendado para reescrituras asistidas por IA)
El servidor MCP es la interfaz principal para usar Agent Shaper con Claude Code. Expone la extracción de formas y la validación de reescrituras como herramientas que el modelo puede llamar directamente.
Agrégalo a tu configuración MCP de Claude Code:
{
"mcpServers": {
"agent-shaper": {
"command": "/path/to/.venv/bin/python",
"args": ["-m", "agent_shaper.mcp_server"]
}
}
}
Herramientas disponibles:
| Herramienta | Propósito |
|---|---|
get_annotated_sources | Ejecutar un script de configuración, trazar el modelo y devolver el código fuente anotado con formas para los archivos solicitados |
get_fx_shapes | Devolver datos FX de formas en bruto para módulos y funciones |
validate_rewrite | Verificar que una clase nn.Module reescrita produce salidas idénticas (atol=1e-5) |
validate_rewrite_function | Verificar que una función independiente reescrita produce salidas idénticas |
validate_file_rewrite | Validar todas las clases en un archivo reescrito en una sola llamada |
save_fixtures | Capturar y persistir fixtures de forward-pass para pruebas posteriores |
generate_test_files | Generar archivos pytest que verifiquen la identidad de salidas contra fixtures guardados |
Flujo de trabajo típico de reescritura:
- Llama a
get_annotated_sourcescon un script de configuración (asignamodel,example_args,dim_namesopcional) y la lista de archivos fuente a anotar. - Lee los fragmentos anotados con formas devueltos para cada clase y función.
- Reescribe el código usando
einsum/einops. - Llama a
validate_rewrite(para clases) ovalidate_rewrite_function(para funciones independientes) para confirmar la identidad numérica antes de escribir en disco.
Contrato del script de configuración — el script debe asignar:
model: una instancia denn.Modulepara trazarexample_args: una tupla de tensores de ejemplo que coincidan con la firma deforward()dim_names(opcional):dictque mapea nombres de dimensiones simbólicas a valores enteros (p. ej.{"B": 4, "T": 16}) para que las formas muestren(B, T)en lugar de(4, 16)
Anotación manual de formas
import torch
from agent_shaper.transformer.model import GPT, GPTConfig
from agent_shaper.fx_utils.manual_annotate import annotate_module_source
cfg = GPTConfig(block_size=32, vocab_size=256, n_layer=2, n_head=2, n_embd=64, dropout=0.0, bias=True)
B, T = 2, 16
example_args = (
torch.randint(0, cfg.vocab_size, (B, T), dtype=torch.long),
torch.randint(0, cfg.vocab_size, (B, T), dtype=torch.long),
)
annotated = annotate_module_source(
GPT(cfg),
example_args,
dim_names={"B": B, "T": T},
output_dir="annotated_output", # writes annotated files here; omit to just get the dict back
)
O ejecuta la prueba de humo integrada directamente:
python -m agent_shaper.fx_utils.manual_annotate
Los archivos de salida se escriben en annotated_output/ preservando la estructura de rutas relativas original.
Anotación LLM
Configura las variables de entorno primero:
export OPENAI_API_KEY=sk-...
export OPENAI_MODEL=gpt-4o
export OPENAI_BASE_URL=https://your-proxy/v1 # optional; omit for default OpenAI
import asyncio, torch
from agent_shaper.transformer.model import GPT, GPTConfig
from agent_shaper.fx_utils.llm_annotate import llm_annotate_module_source, AnnotationMode
cfg = GPTConfig(block_size=32, vocab_size=256, n_layer=2, n_head=2, n_embd=64, dropout=0.0, bias=True)
B, T = 2, 16
example_args = (
torch.randint(0, cfg.vocab_size, (B, T), dtype=torch.long),
torch.randint(0, cfg.vocab_size, (B, T), dtype=torch.long),
)
async def main():
annotated = await llm_annotate_module_source(
GPT(cfg),
example_args,
mode=AnnotationMode.COMMENT, # or AnnotationMode.EINSUM
dim_names={"B": B, "T": T},
output_dir="llm_annotated_output",
)
asyncio.run(main())
O ejecuta la prueba de humo integrada:
python -m agent_shaper.fx_utils.llm_annotate
Revisión de diffs
Por defecto, después de generar el archivo anotado por LLM, VS Code se abre automáticamente mostrando un diff lado a lado del original vs. el archivo reescrito (open_in_vscode=True). Pasa open_in_vscode=False para suprimir esto.
También puedes usar el visor de diffs de Streamlit para una revisión en el navegador:
.venv/bin/streamlit run agent_shaper/fx_utils/diff_viewer.py
Ingresa la ruta al archivo original a la izquierda y el archivo generado (p. ej. llm_annotated_output/agent_shaper/transformer/model.py) a la derecha. El visor renderiza un diff unificado con resaltado de sintaxis.
dim_names
El parámetro opcional dim_names mapea nombres simbólicos a sus valores concretos en la ejecución de ejemplo. Esto permite que las anotaciones de forma muestren (B, T, n_embd) en lugar de (2, 16, 64). Cuando dos nombres comparten el mismo valor (p. ej. B=2 y n_head=2), la anotación muestra B/n_head.
dim_names = {"B": 2, "T": 16, "n_embd": 64, "n_head": 2}
get_module_shapes directamente
from agent_shaper.fx_utils import get_module_shapes, TensorInfo
module_infos = get_module_shapes(model, example_args, dim_names=dim_names)
for info in module_infos:
print(info.class_name, info.source_file, info.line_start, info.line_end)
for t in info.tensors:
print(" ", t.name, t.shape, t.annotated_shape)
Cada ModuleInfo contiene:
class_name,module_origin,source_file,line_start,line_endparameters— lista deTensorInfopara entradasnn.Parameterde__init__tensors— lista deTensorInfopara cada nodo FX intermedio en el forward pass
Las instancias de módulos repetidas con secuencias de formas idénticas (p. ej. bloques de transformadores) se deduplican a una sola entrada.
Ejemplos
La carpeta examples/ contiene reescrituras trabajadas en varias familias de modelos, cada una emparejada con una versión einsum/einops validada:
examples/transformer/— GPT-2, LLaMA, Qwen3, Swin Transformer. Cada modelo tiene una contraparte*_einsum.pyreescrita contorch.einsumyeinops.examples/alignment/— funciones independientes de pérdida de alineación (DPO, IPO, SimPO, JSD, OPD) tanto en forma original como einsum.examples/direct_alignment/— implementaciones de pérdida RLHF estilo producción (DPO, cDPO, IPO, SimPO, ORPO, KTO, APO-zero, APO-down) como clasesnn.Modulecon un bucle de entrenamiento completo. El archivoloss.pyusaeinops.einsumyeinops.reduceen todo momento, validado contra los originales basados en gather con atol=1e-5.
Todas las reescrituras *_einsum.py fueron validadas usando las herramientas MCP validate_rewrite / validate_rewrite_function.