PyTorch-FX Shaper

O servidor MCP fornece a forma dos tensores para converter código PyTorch para einsum e einops.

Documentação

Agent Shaper

O Agent Shaper extrai metadados de formato de tensores por módulo de qualquer PyTorch nn.Module e os utiliza para anotar arquivos de código-fonte — seja com comentários descritivos de formato ou reescrevendo operações como torch.einsum / einops.

Como funciona

  1. Extração de formatos (fx_utils/get_fx_data.py) — executa torch.export + ShapeProp no seu módulo para capturar o formato de cada tensor intermediário no forward pass de cada módulo definido no workspace, não apenas as entradas.
  2. Anotação manual (fx_utils/manual_annotate.py) — insere os formatos extraídos como comentários inline no final de cada linha de código relevante.
  3. Anotação por LLM (fx_utils/llm_annotate.py) — alimenta cada classe de módulo anotada manualmente a um LLM em paralelo. Dois modos:
    • COMMENT — reescreve comentários com nomes descritivos de dimensões (ex.: batch_size, seq_len, n_embd) e explicações em linguagem simples de cada transformação.
    • EINSUM — reescreve o módulo inteiro substituindo multiplicações de matrizes e operações de atenção por torch.einsum, colapsando redimensionamentos intermediários quando possível.
  4. Visualizador de diffs — revise as alterações do LLM antes de aceitá-las, seja em uma interface Streamlit ou diretamente no editor de diff nativo do VS Code.
  5. Servidor MCP (agent_shaper/mcp_server.py) — expõe todo o pipeline de extração de formatos e validação de reescrita como ferramentas MCP consumíveis pelo Claude Code ou qualquer assistente de IA compatível com MCP.

Todos os módulos em um arquivo são processados em paralelo via asyncio. Tudo fora das classes de módulos (imports, dataclasses, objetos de configuração) é preservado inalterado no arquivo de saída.

Estrutura do projeto

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 arquivo *_einsum.py é a contraparte reescrita com einsum/einops do original, validado para produzir saídas numericamente idênticas (atol=1e-5).

Instalação

Requer Python 3.12+.

python -m venv .venv
source .venv/bin/activate
pip install -e .

Uso

Servidor MCP (recomendado para reescritas assistidas por IA)

O servidor MCP é a interface principal para usar o Agent Shaper com o Claude Code. Ele expõe extração de formatos e validação de reescrita como ferramentas que o modelo pode chamar diretamente.

Adicione-o à configuração MCP do seu Claude Code:

{
  "mcpServers": {
    "agent-shaper": {
      "command": "/path/to/.venv/bin/python",
      "args": ["-m", "agent_shaper.mcp_server"]
    }
  }
}

Ferramentas disponíveis:

FerramentaFinalidade
get_annotated_sourcesExecuta um script de configuração, rastreia o modelo e retorna o código-fonte anotado com formatos para os arquivos solicitados
get_fx_shapesRetorna dados brutos de formatos FX para módulos e funções
validate_rewriteVerifica se uma classe nn.Module reescrita produz saídas idênticas (atol=1e-5)
validate_rewrite_functionVerifica se uma função independente reescrita produz saídas idênticas
validate_file_rewriteValida todas as classes em um arquivo reescrito em uma única chamada
save_fixturesCaptura e persiste fixtures de forward-pass para testes posteriores
generate_test_filesGera arquivos pytest que verificam identidade de saída contra fixtures salvos

Fluxo de trabalho típico de reescrita:

  1. Chame get_annotated_sources com um script de configuração (atribui model, example_args, dim_names opcional) e a lista de arquivos de código-fonte a anotar.
  2. Leia os trechos anotados com formatos retornados para cada classe e função.
  3. Reescreva o código usando einsum / einops.
  4. Chame validate_rewrite (para classes) ou validate_rewrite_function (para funções independentes) para confirmar identidade numérica antes de gravar em disco.

Contrato do script de configuração — o script deve atribuir:

  • model: uma instância de nn.Module para rastrear
  • example_args: uma tupla de tensores de exemplo correspondentes à assinatura de forward()
  • dim_names (opcional): dict mapeando nomes simbólicos de dimensões para valores inteiros (ex.: {"B": 4, "T": 16}) para que os formatos mostrem (B, T) em vez de (4, 16)

Anotação manual de formatos

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
)

Ou execute o teste de fumaça integrado diretamente:

python -m agent_shaper.fx_utils.manual_annotate

Os arquivos de saída são gravados em annotated_output/ preservando a estrutura de caminhos relativos original.

Anotação por LLM

Defina as variáveis de ambiente primeiro:

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())

Ou execute o teste de fumaça integrado:

python -m agent_shaper.fx_utils.llm_annotate

Revisando diffs

Por padrão, após gerar o arquivo anotado pelo LLM, o VS Code abre automaticamente mostrando um diff lado a lado do original vs. o arquivo reescrito (open_in_vscode=True). Passe open_in_vscode=False para suprimir isso.

Você também pode usar o visualizador de diffs Streamlit para uma revisão baseada em navegador:

.venv/bin/streamlit run agent_shaper/fx_utils/diff_viewer.py

Digite o caminho do arquivo original à esquerda e o arquivo gerado (ex.: llm_annotated_output/agent_shaper/transformer/model.py) à direita. O visualizador renderiza um diff unificado com realce de sintaxe.

dim_names

O parâmetro opcional dim_names mapeia nomes simbólicos para seus valores concretos na execução de exemplo. Isso permite que as anotações de formato mostrem (B, T, n_embd) em vez de (2, 16, 64). Quando dois nomes compartilham o mesmo valor (ex.: B=2 e n_head=2), a anotação mostra B/n_head.

dim_names = {"B": 2, "T": 16, "n_embd": 64, "n_head": 2}

get_module_shapes diretamente

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 contém:

  • 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 nó FX intermediário no forward pass

Instâncias de módulos repetidas com sequências de formatos idênticas (ex.: blocos de transformadores) são deduplicadas para uma única entrada.

Exemplos

A pasta examples/ contém reescritas trabalhadas em várias famílias de modelos, cada uma emparelhada com uma versão einsum/einops validada:

  • examples/transformer/ — GPT-2, LLaMA, Qwen3, Swin Transformer. Cada modelo tem uma contraparte *_einsum.py reescrita com torch.einsum e einops.
  • examples/alignment/ — funções independentes de perda de alinhamento (DPO, IPO, SimPO, JSD, OPD) tanto na forma original quanto na forma einsum.
  • examples/direct_alignment/ — implementações de perda RLHF em estilo de produção (DPO, cDPO, IPO, SimPO, ORPO, KTO, APO-zero, APO-down) como classes nn.Module com um loop de treinamento completo. O arquivo loss.py usa einops.einsum e einops.reduce em todo o código, validado contra os originais baseados em gather com atol=1e-5.

Todas as reescritas *_einsum.py foram validadas usando as ferramentas MCP validate_rewrite / validate_rewrite_function.