import numpy as np
import torch
from typing import Optional
from .tonnetz import distance_matrix
class ToroidalMask:
def __init__(self, seq_len: int, radius: float = 2.0, alpha: float = 1.0,
grid_size: int = 12, mask_type: str = "hybrid"):
self.seq_len = seq_len
self.radius = radius
self.alpha = alpha
self.grid_size = grid_size
if mask_type not in ("hybrid", "hard_cutoff", "soft_exponential"):
raise ValueError(f"Unknown mask_type: {mask_type}")
self.mask_type = mask_type
@classmethod
def hard_cutoff(cls, seq_len: int, radius: float, grid_size: int = 12) -> "ToroidalMask":
return cls(seq_len, radius=radius, alpha=0.0, grid_size=grid_size,
mask_type="hard_cutoff")
@classmethod
def soft_exponential(cls, seq_len: int, alpha: float, grid_size: int = 12) -> "ToroidalMask":
return cls(seq_len, radius=0.0, alpha=alpha, grid_size=grid_size,
mask_type="soft_exponential")
@classmethod
def hybrid(cls, seq_len: int, radius: float = 2.0, alpha: float = 1.0,
grid_size: int = 12) -> "ToroidalMask":
return cls(seq_len, radius=radius, alpha=alpha, grid_size=grid_size,
mask_type="hybrid")
def generate(self) -> np.ndarray:
dist = distance_matrix(self.seq_len, self.grid_size).astype(np.float32)
if self.mask_type == "hard_cutoff":
return np.where(dist <= self.radius, 1.0, 0.0).astype(np.float32)
elif self.mask_type == "soft_exponential":
return np.exp(-self.alpha * dist).astype(np.float32)
else: return np.where(
dist <= self.radius,
1.0,
np.exp(-self.alpha * (dist - self.radius))
).astype(np.float32)
def to_tensor(self, device: Optional[torch.device] = None) -> torch.Tensor:
mask_np = self.generate()
t = torch.from_numpy(mask_np)
if device is not None:
t = t.to(device)
return t
class SparseMask:
def __init__(self, mask: ToroidalMask, threshold: float = 0.01):
dense = mask.generate()
rows, cols = np.nonzero(dense > threshold)
values = dense[rows, cols]
self.size = mask.seq_len
self._indices = torch.stack([
torch.from_numpy(rows.astype(np.int64)),
torch.from_numpy(cols.astype(np.int64)),
])
self._values = torch.from_numpy(values)
self._sparse = torch.sparse_coo_tensor(
self._indices, self._values, (self.size, self.size)
)
@property
def nnz(self) -> int:
return self._values.shape[0]
@property
def sparsity(self) -> float:
total = self.size * self.size
return 1.0 - (self.nnz / total) if total > 0 else 0.0
def to_dense(self) -> torch.Tensor:
return self._sparse.to_dense()
def sinkhorn_knopp(matrix: torch.Tensor, n_iters: int = 20) -> torch.Tensor:
M = torch.exp(matrix) if matrix.min() < 0 else matrix.clone().float()
M = M + 1e-8 for _ in range(n_iters):
M = M / M.sum(dim=-1, keepdim=True)
M = M / M.sum(dim=-2, keepdim=True)
return M
def is_doubly_stochastic(matrix: torch.Tensor, tol: float = 0.01) -> bool:
row_sums = matrix.sum(dim=-1)
col_sums = matrix.sum(dim=-2)
return (
torch.all(torch.abs(row_sums - 1.0) < tol).item()
and torch.all(torch.abs(col_sums - 1.0) < tol).item()
)