Download src/bigru_t/utils/tensor_ops.py from PowerMachine/BiGRU_T_version: direct link, hf CLI and curl.
- Browser
- Download file 2.18 kB
-
https://huggingface.co/PowerMachine/BiGRU_T_version/resolve/main/src/bigru_t/utils/tensor_ops.py
- Command line
-
hf download hf://PowerMachine/BiGRU_T_version/src/bigru_t/utils/tensor_ops.py
-
curl -L -o tensor_ops.py https://huggingface.co/PowerMachine/BiGRU_T_version/resolve/main/src/bigru_t/utils/tensor_ops.py
2.18 kB
| """Xavante - tensor_ops.py - Operacoes tensoriais otimizadas.""" | |
| from __future__ import annotations | |
| import logging | |
| import math | |
| from typing import Tuple | |
| import torch | |
| logger = logging.getLogger(__name__) | |
| def safe_softmax(x: torch.Tensor, dim: int = -1) -> torch.Tensor: | |
| """Softmax com subtracao do max para estabilidade numerica.""" | |
| return torch.softmax(x - x.amax(dim=dim, keepdim=True), dim=dim) | |
| def safe_log(x: torch.Tensor, eps: float = 1e-9) -> torch.Tensor: | |
| return torch.log(x.clamp_min(eps)) | |
| def cosine_similarity(a: torch.Tensor, b: torch.Tensor, dim: int = -1, eps: float = 1e-8) -> torch.Tensor: | |
| return torch.nn.functional.cosine_similarity(a, b, dim=dim, eps=eps) | |
| def orthogonal_init(weight: torch.Tensor) -> None: | |
| """Inicializa weight como matriz ortogonal.""" | |
| nn_init = torch.nn.init | |
| nn_init.orthogonal_(weight) | |
| def skew_to_rotation(skew: torch.Tensor) -> torch.Tensor: | |
| """Converte matriz skew-symmetric para rotation via Rodrigues.""" | |
| n = skew.shape[0] | |
| I = torch.eye(n, device=skew.device, dtype=skew.dtype) | |
| A = skew | |
| A2 = A @ A | |
| # Exponential via series truncada | |
| return I + A + 0.5 * A2 + (1.0 / 6.0) * (A2 @ A) | |
| def chunked_matmul(a: torch.Tensor, b: torch.Tensor, chunk: int = 4096) -> torch.Tensor: | |
| """Matmul em chunks para evitar OOM com matrizes grandes.""" | |
| if a.shape[1] <= chunk: | |
| return a @ b | |
| out_chunks = [] | |
| for i in range(0, a.shape[0], chunk): | |
| out_chunks.append(a[i : i + chunk] @ b) | |
| return torch.cat(out_chunks, dim=0) | |
| def padding_mask(lengths: torch.Tensor, max_len: int) -> torch.Tensor: | |
| """Cria mascara bool [B, max_len] onde True = token valido.""" | |
| idx = torch.arange(max_len, device=lengths.device).unsqueeze(0) | |
| return idx < lengths.unsqueeze(1) | |
| def causal_mask(seq_len: int, device: torch.device) -> torch.Tensor: | |
| """Mascara causal [seq_len, seq_len] (1 = permitido).""" | |
| return torch.tril(torch.ones(seq_len, seq_len, device=device)) | |
| __all__ = [ | |
| "safe_softmax", | |
| "safe_log", | |
| "cosine_similarity", | |
| "orthogonal_init", | |
| "skew_to_rotation", | |
| "chunked_matmul", | |
| "padding_mask", | |
| "causal_mask", | |
| ] | |