Skip to content

Commit d29f868

Browse files
committed
fix: elem kernel extra N arg; variance nested glu_type
- Remove extra N arg in _softsign_glu_bwd_elem_kernel call (M,N,N → M,N) - Read variance glu_type from nested predictor config paths - Assert both predictors use softsign_glu before patching
1 parent 9256c0b commit d29f868

2 files changed

Lines changed: 8 additions & 6 deletions

File tree

modules/kernels/fused_linear_softsign_glu.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -337,7 +337,7 @@ def elem_grid(meta):
337337
_softsign_glu_bwd_elem_kernel[elem_grid](
338338
left, gate, grad_y,
339339
grad_left_pre, grad_gate,
340-
M, N, N,
340+
M, N,
341341
left.stride(0), left.stride(1),
342342
gate.stride(0), gate.stride(1),
343343
grad_y.stride(0), grad_y.stride(1),

training/variance_task.py

Lines changed: 7 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -124,11 +124,13 @@ def __init__(self):
124124
if hparams.get('use_fused_kernels', False):
125125
from modules.kernels.integration import patch_variance_model
126126
from lightning.pytorch.utilities.rank_zero import rank_zero_info
127-
n = patch_variance_model(
128-
self.model,
129-
glu_type=hparams.get('backbone_args', {}).get('glu_type', 'softsign_glu'),
130-
)
131-
rank_zero_info('Fused kernels: patched %d LYNXNet2 blocks in variance model', n)
127+
# Read glu_type from nested predictor configs
128+
pitch_glu = hparams.get('pitch_prediction_args', {}).get('backbone_args', {}).get('glu_type', 'softsign_glu')
129+
var_glu = hparams.get('variances_prediction_args', {}).get('backbone_args', {}).get('glu_type', 'softsign_glu')
130+
assert pitch_glu == var_glu == 'softsign_glu', \
131+
f"Fused kernels only support softsign_glu, got pitch={pitch_glu} var={var_glu}"
132+
n = patch_variance_model(self.model, glu_type='softsign_glu')
133+
rank_zero_info('Fused kernels: patched %d LYNXNet2 blocks in variance model (softsign_glu)', n)
132134

133135

134136
def _build_model(self):

0 commit comments

Comments
 (0)