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
- Extração de formatos (
fx_utils/get_fx_data.py) — executatorch.export+ShapePropno 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. - 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. - 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.
- COMMENT — reescreve comentários com nomes descritivos de dimensões (ex.:
- 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.
- 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:
| Ferramenta | Finalidade |
|---|---|
get_annotated_sources | Executa um script de configuração, rastreia o modelo e retorna o código-fonte anotado com formatos para os arquivos solicitados |
get_fx_shapes | Retorna dados brutos de formatos FX para módulos e funções |
validate_rewrite | Verifica se uma classe nn.Module reescrita produz saídas idênticas (atol=1e-5) |
validate_rewrite_function | Verifica se uma função independente reescrita produz saídas idênticas |
validate_file_rewrite | Valida todas as classes em um arquivo reescrito em uma única chamada |
save_fixtures | Captura e persiste fixtures de forward-pass para testes posteriores |
generate_test_files | Gera arquivos pytest que verificam identidade de saída contra fixtures salvos |
Fluxo de trabalho típico de reescrita:
- Chame
get_annotated_sourcescom um script de configuração (atribuimodel,example_args,dim_namesopcional) e a lista de arquivos de código-fonte a anotar. - Leia os trechos anotados com formatos retornados para cada classe e função.
- Reescreva o código usando
einsum/einops. - Chame
validate_rewrite(para classes) ouvalidate_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 denn.Modulepara rastrearexample_args: uma tupla de tensores de exemplo correspondentes à assinatura deforward()dim_names(opcional):dictmapeando 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_endparameters— lista deTensorInfopara entradasnn.Parameterde__init__tensors— lista deTensorInfopara 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.pyreescrita comtorch.einsumeeinops.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 classesnn.Modulecom um loop de treinamento completo. O arquivoloss.pyusaeinops.einsumeeinops.reduceem 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.