BiGRU_T_version / src /bigru_t /training /gradient_surgery.py
PowerMachine's picture
Upload folder using huggingface_hub
3275441 verified
Raw History Blame Contribute Delete
4.03 kB
"""gradient_surgery.py — Lema 2: gradiente cirúrgico (projeção ortogonal).
Implementa o Lema 2 (Gradiente cirúrgico).
Quando os gradientes da perda principal (loss_main) e da perda de hipótese
(loss_hyp) são conflitantes (produto interno negativo), projeta o gradiente
de loss_hyp no complemento ortogonal de loss_main. Isso garante que ambas
as perdas possam ser reduzidas sem interferência destrutiva.
Reaproveita a filosofia do punishment_reward_system.py do v13.9.2, mas:
- Substitui o gate binário de punição por projeção ortogonal contínua
- Não há "patience" nem "skip steps" — a cirurgia é aplicada sempre que
loss_hyp é calculada
- O limiar tau (do Lema 4) decide SE loss_hyp é calculada, não COMO o
gradiente é combinado
Referências:
- Yu et al. (2020). Gradient Surgery for Multi-Task Learning. NeurIPS.
- Csiszár, P. (1984). Information Geometry and Orthogonal Projections.
"""
from __future__ import annotations
from typing import Iterable
import torch
import torch.nn as nn
def orthogonalize_gradient(
g_main: torch.Tensor, g_hyp: torch.Tensor, eps: float = 1e-8
) -> torch.Tensor:
"""Projeta g_hyp no complemento ortogonal de g_main.
Se <g_main, g_hyp> < 0 (conflitantes):
g_hyp_perp = g_hyp - (<g_main, g_hyp> / ||g_main||^2) * g_main
Caso contrário (alinhados):
g_hyp_perp = g_hyp (sem projeção)
Args:
g_main: gradiente da perda principal (tensor)
g_hyp: gradiente da perda de hipótese (tensor, mesmo shape)
eps: estabilidade numérica
Returns:
g_hyp_perp: gradiente projetado (mesmo shape)
"""
# Produto interno escalar (flatten para generalidade)
g_main_flat = g_main.flatten()
g_hyp_flat = g_hyp.flatten()
dot = torch.dot(g_main_flat, g_hyp_flat)
if dot < 0:
norm_sq = torch.dot(g_main_flat, g_main_flat)
if norm_sq > eps:
# g_hyp_perp = g_hyp - (dot / ||g_main||^2) * g_main
g_hyp_perp = g_hyp - (dot / norm_sq) * g_main
else:
# g_main ≈ 0: não há direção para projetar
g_hyp_perp = g_hyp
else:
# Gradientes alinhados: sem projeção
g_hyp_perp = g_hyp
return g_hyp_perp
def apply_gradient_surgery(
model: nn.Module,
loss_main: torch.Tensor,
loss_hyp: torch.Tensor,
) -> None:
"""Aplica cirurgia de gradiente e atribui p.grad combinado.
Calcula gradientes independentes para loss_main e loss_hyp, aplica
projeção ortogonal quando conflitantes, e atribui o gradiente combinado
a p.grad para cada parâmetro p do modelo.
Args:
model: o modelo (parâmetros serão atualizados in-place via .grad)
loss_main: perda principal (escalar, com grad)
loss_hyp: perda de hipótese (escalar, com grad)
Nota:
- Esta função NÃO chama optimizer.step() — apenas atribui p.grad.
- O chamador deve fazer optimizer.step() e optimizer.zero_grad().
- loss_main e loss_hyp devem compartilhar o grafo computacional
(ou loss_hyp deve ser computada com retain_graph=True).
"""
# Zera gradientes existentes
model.zero_grad(set_to_none=True)
# Gradientes da perda principal (retain_graph=True para reusar na hyp)
params = [p for p in model.parameters() if p.requires_grad]
grads_main = torch.autograd.grad(
loss_main, params, retain_graph=True, create_graph=False, allow_unused=True
)
# Gradientes da perda de hipótese
grads_hyp = torch.autograd.grad(
loss_hyp, params, retain_graph=False, create_graph=False, allow_unused=True
)
# Combina com cirurgia ortogonal
for p, g_m, g_h in zip(params, grads_main, grads_hyp):
if g_m is None:
g_m = torch.zeros_like(p)
if g_h is None:
g_h = torch.zeros_like(p)
# Cirurgia ortogonal (Lema 2)
g_h_orth = orthogonalize_gradient(g_m, g_h)
# Gradiente combinado
p.grad = (g_m + g_h_orth).detach()
return