Download src/bigru_t/training/gradient_surgery.py from PowerMachine/BiGRU_T_version: direct link, hf CLI and curl.
- Browser
- Download file 4.03 kB
-
https://huggingface.co/PowerMachine/BiGRU_T_version/resolve/main/src/bigru_t/training/gradient_surgery.py
- Command line
-
hf download hf://PowerMachine/BiGRU_T_version/src/bigru_t/training/gradient_surgery.py
-
curl -L -o gradient_surgery.py https://huggingface.co/PowerMachine/BiGRU_T_version/resolve/main/src/bigru_t/training/gradient_surgery.py
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 | |