Compreendendo torch.outer no PyTorch: Uso, migração e compatibilidade entre versões

No ecossistema do PyTorch, o cálculo do produto externo entre dois vetores é uma operação comum em tarefas como processamento de sinais, implementações personalizadas de mecanismos de atenção ou cruzamento de características. Historicamente, essa operação era relaizada com a função torch.ger. No entanto, a partir da versão 1.7.0 do PyTorch, essa função foi substituída por torch.outer, que oferece a mesma funcionalidade com um nome mais intuitivo.

O produto externo entre dois vetores unidimensionais a (de comprimento n) e b (de comprimento m) resulta em uma matriz de dimensão n × m, onde cada elemento na posição (i, j) é dado por a[i] * b[j].

Uso básico de torch.outer

A assinatura da função é simples:

torch.outer(input, vec2, *, out=None)

Exemplo prático:

import torch

x = torch.tensor([2, 3, 4, 5])
y = torch.tensor([1, 2, 3])

result = torch.outer(x, y)
print(result)
# Saída:
# tensor([[ 2,  4,  6],
#         [ 3,  6,  9],
#         [ 4,  8, 12],
#         [ 5, 10, 15]])
print(result.shape)  # torch.Size([4, 3])

Esta operação é matematicamente equivalente ao antigo torch.ger. De fato, testes mostram que ambos produzem resultados idênticos bit a bit quando disponíveis:

v1 = torch.randn(50)
v2 = torch.randn(30)

assert torch.equal(torch.outer(v1, v2), torch.ger(v1, v2))

Compatibilidade entre versões do PyTorch

A transição entre as funções está diretamente ligada à versão do PyTorch utilizada:

  • Versões ≤ 1.6.0: Apenas torch.ger está disponível.
  • Versão 1.7.0: Ambas as funções coexistem; torch.ger é marcado como obsoleto.
  • Versões ≥ 1.8.0: Apenas torch.outer está disponível; torch.ger foi removido.

Erros comuns incluem AttributeError ao tentar usar torch.outer em versões antigas ou torch.ger em versões recentes.

Estratégias para código compatível com múltiplas versões

Para garantir que seu código funcione independentemente da versão do PyTorch, recomenda-se detectar dinamicamente qual função está disponível:

import torch
import warnings

def compute_outer(u, v):
    if hasattr(torch, 'outer'):
        return torch.outer(u, v)
    elif hasattr(torch, 'ger'):
        warnings.warn(
            "torch.ger is deprecated. Use torch.outer (PyTorch ≥ 1.7).",
            DeprecationWarning,
            stacklevel=2
        )
        return torch.ger(u, v)
    else:
        raise RuntimeError("Nenhuma função de produto externo disponível.")

Alternativamente, crie um módulo de compatibilidade centralizado:

# compat.py
import torch

outer = torch.outer if hasattr(torch, 'outer') else torch.ger

E use em todo o projeto:

from .compat import outer

z = outer(tensor_a, tensor_b)

Boas práticas de gerenciamento de dependências

Além de adaptar o código, defina claramente os requisitos de versão no arquivo de dependências:

# requirements.txt
torch>=1.8.0  # Se não for necessário suportar versões antigas

Use ambientes virtuais (com venv ou conda) e, em produção, considere o uso de contêineres (como Docker) para garantir consistência entre ambientes de desenvolvimento e execução.

Tags: Pytorch tensor outer-product api-compatibility version-migration

Publicado em 10-7 22:38