"""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 < 0 (conflitantes): g_hyp_perp = 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