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

  1. Extracción de formas (fx_utils/get_fx_data.py) — ejecuta torch.export + ShapeProp en 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.
  2. 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.
  3. 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.
  4. 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.
  5. 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:

HerramientaPropósito
get_annotated_sourcesEjecutar un script de configuración, trazar el modelo y devolver el código fuente anotado con formas para los archivos solicitados
get_fx_shapesDevolver datos FX de formas en bruto para módulos y funciones
validate_rewriteVerificar que una clase nn.Module reescrita produce salidas idénticas (atol=1e-5)
validate_rewrite_functionVerificar que una función independiente reescrita produce salidas idénticas
validate_file_rewriteValidar todas las clases en un archivo reescrito en una sola llamada
save_fixturesCapturar y persistir fixtures de forward-pass para pruebas posteriores
generate_test_filesGenerar archivos pytest que verifiquen la identidad de salidas contra fixtures guardados

Flujo de trabajo típico de reescritura:

  1. Llama a get_annotated_sources con un script de configuración (asigna model, example_args, dim_names opcional) y la lista de archivos fuente a anotar.
  2. Lee los fragmentos anotados con formas devueltos para cada clase y función.
  3. Reescribe el código usando einsum / einops.
  4. Llama a validate_rewrite (para clases) o validate_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 de nn.Module para trazar
  • example_args: una tupla de tensores de ejemplo que coincidan con la firma de forward()
  • dim_names (opcional): dict que 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_end
  • parameters — lista de TensorInfo para entradas nn.Parameter de __init__
  • tensors — lista de TensorInfo para 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.py reescrita con torch.einsum y einops.
  • 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 clases nn.Module con un bucle de entrenamiento completo. El archivo loss.py usa einops.einsum y einops.reduce en 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.