Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
22 commits
Select commit Hold shift + click to select a range
d8b261c
feat: Triton fused SoftSignGLU kernel for LYNXNet2
KakaruHayate Jul 2, 2026
e357c71
perf: skip variance fusion, realistic warmup M, fewer autotune configs
KakaruHayate Jul 2, 2026
99dad24
fix: remove useless warmup from __init__ (model on CPU, no compilation)
KakaruHayate Jul 2, 2026
06cc355
fix: disable variance fusion by default (small backbones, no benefit)
KakaruHayate Jul 2, 2026
6058ec5
fix: remove useless warmup from __init__ (model on CPU, no compilation)
KakaruHayate Jul 2, 2026
9256c0b
fix: address PR review issues
KakaruHayate Jul 2, 2026
d29f868
fix: elem kernel extra N arg; variance nested glu_type
KakaruHayate Jul 2, 2026
ade184c
fix: backward kernel dtype not hardcoded fp16
KakaruHayate Jul 3, 2026
cf25e36
fix: optimize Triton fused SoftSignGLU kernel for training
KakaruHayate Jul 7, 2026
bce80cd
fix: warn when fused kernels are skipped due to unsupported glu_type
KakaruHayate Jul 7, 2026
b336d78
review: fix hard blocks — TF32 removal, glu_type revert, robustness
KakaruHayate Jul 14, 2026
0cafbbe
review round 2: variance warmup + GPU-verified test rework
KakaruHayate Jul 15, 2026
407b315
fix: rank_zero_info unbound in acoustic ImportError path
KakaruHayate Jul 15, 2026
829f88b
chore: add missing trailing newline in integration.py
KakaruHayate Jul 15, 2026
9801e5d
feat: DoubleSoftSignGLU (FastWaveD-style) + ATanGLUFunction-style mem…
KakaruHayate Jul 15, 2026
9a4d60c
refactor: unify use_fused_kernels_variance into use_fused_kernels
KakaruHayate Jul 15, 2026
f8056c2
fix: fail fast when Triton is unavailable
KakaruHayate Jul 21, 2026
c13588f
fix: warm only patched variance backbones
KakaruHayate Aug 3, 2026
222a72b
fix: address CodeRabbit fused kernel review
KakaruHayate Aug 6, 2026
6133bb5
fix: fail fast when fused kernels lack Triton
KakaruHayate Aug 6, 2026
8cd54bf
refactor: finalize fused SoftSignGLU integration
KakaruHayate Aug 6, 2026
6716397
fix: isolate warmup RNG, restore 1D input rank
KakaruHayate Aug 7, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions configs/acoustic.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,9 @@ base_config:

task_cls: training.acoustic_task.AcousticTask

# Enable Triton-fused Linear+SoftSignGLU kernels for LYNXNet2 backbones.
use_fused_kernels: false

dictionaries: {}
extra_phonemes: []
merged_phoneme_groups: []
Expand Down
3 changes: 3 additions & 0 deletions configs/variance.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,9 @@ base_config:

task_cls: training.variance_task.VarianceTask

# Enable Triton-fused Linear+SoftSignGLU kernels for LYNXNet2 backbones.
use_fused_kernels: false

dictionaries: {}
extra_phonemes: []
merged_phoneme_groups: []
Expand Down
6 changes: 5 additions & 1 deletion modules/backbones/lynxnet2.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,9 @@
import torch.nn as nn
import torch.nn.functional as F

from modules.commons.common_layers import SinusoidalPosEmb, SwiGLU, ATanGLU, Transpose, AdamWLinear
from modules.commons.common_layers import (
SinusoidalPosEmb, SwiGLU, ATanGLU, SoftSignGLU, Transpose, AdamWLinear
)
from utils.hparams import hparams


Expand All @@ -14,6 +16,8 @@ def __init__(self, dim, expansion_factor, kernel_size=31, dropout=0., glu_type='
_glu = SwiGLU()
elif glu_type == 'atanglu':
_glu = ATanGLU()
elif glu_type == 'softsign_glu':
_glu = SoftSignGLU()
else:
raise ValueError(f'{glu_type} is not a valid activation')
if float(dropout) > 0.:
Expand Down
41 changes: 41 additions & 0 deletions modules/commons/common_layers.py
Original file line number Diff line number Diff line change
Expand Up @@ -175,6 +175,47 @@ def forward(self, x):
return out * torch.atan(gate)


class SoftSignGLUFunction(torch.autograd.Function):
"""ATanGLUFunction-style memory trick for SoftSignGLU.

softsign'(x) = 1/(1+|x|)^2 = (1-|softsign(x)|)^2, so both partial
derivatives of y = out * softsign(gate) are precomputable in forward:
dy/dout = softsign(gate)
dy/dgate = out * (1-|softsign(gate)|)^2
Saves 2 tensors (vs 3 for naive autograd) and backward is two pure
multiplies with no softsign recompute.
"""
@staticmethod
def forward(ctx, out, gate):
ss_gate = torch.nn.functional.softsign(gate)
decay_out = out * (1.0 - ss_gate.abs()).square()
ctx.save_for_backward(ss_gate, decay_out)
return out * ss_gate

@staticmethod
def backward(ctx, grad_output):
ss_gate, decay_out = ctx.saved_tensors
return grad_output * ss_gate, grad_output * decay_out


class SoftSignGLU(nn.Module):
"""Gated Linear Unit with SoftSign gate: out * softsign(gate).

More numerically stable than ATanGLU (no approximation needed in
Triton kernels) while providing similar gating behavior.
"""
def __init__(self, dim=-1):
super().__init__()
self.dim = dim

def forward(self, x):
out, gate = torch.split(x, x.size(self.dim) // 2, dim=self.dim)
if self.training:
return SoftSignGLUFunction.apply(out, gate)
else:
return out * torch.nn.functional.softsign(gate)


class AdamWConv1d(torch.nn.Conv1d):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
Expand Down
1 change: 1 addition & 0 deletions modules/kernels/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
# Fused kernels for LYNXNet2 optimization
Loading