s2pro-egy / scripts /patch_audio_decoder.py
ehabnegm's picture
Phase 1: merged fast-AR LoRA step-1200 + full toolkit + phase-2 handoff
5c2beba verified
Raw
History Blame Contribute Delete
3.46 kB
#!/usr/bin/env python3
"""Patch sglang-omni s2-pro audio_decoder: SDPA fallback for flash_attn_with_kvcache.
FA3 (sgl_kernel.flash_attn) has no sm_120 kernels; the Fast-AR attends over
<=11 positions so plain SDPA is numerically fine and fast. Env-gated:
FISH_FORCE_SDPA=1 activates the fallback (default keeps upstream FA3 path).
"""
from pathlib import Path
P = Path("/opt/work/sglang-omni/sglang_omni/models/fishaudio_s2_pro/"
"fish_speech/models/text2semantic/audio_decoder.py")
src = P.read_text()
anchor = ''') -> torch.Tensor:
return flash_attn_with_kvcache(
q=q,
k_cache=k_cache,
v_cache=v_cache,
k=k,
v=v,
cache_seqlens=cache_seqlens.contiguous() if cache_seqlens is not None else None,
causal=causal,
num_splits=num_splits,
)'''
replacement = ''') -> torch.Tensor:
if FISH_FORCE_SDPA:
return _sdpa_attn_with_kvcache(q, k_cache, v_cache, k, v, cache_seqlens, causal)
return flash_attn_with_kvcache(
q=q,
k_cache=k_cache,
v_cache=v_cache,
k=k,
v=v,
cache_seqlens=cache_seqlens.contiguous() if cache_seqlens is not None else None,
causal=causal,
num_splits=num_splits,
)
def _sdpa_attn_with_kvcache(q, k_cache, v_cache, k, v, cache_seqlens, causal):
"""Pure-torch drop-in for flash_attn_with_kvcache (sm_120-safe).
q: (b, sq, hq, d); k_cache/v_cache: (b, T, hk, d) mutated in place;
k/v: (b, sn, hk, d) appended at cache_seqlens; returns (b, sq, hq, d).
"""
b, sq, hq, d = q.shape
hk = k_cache.shape[2]
if cache_seqlens is None:
cache_seqlens = torch.zeros(b, dtype=torch.int32, device=q.device)
cache_seqlens = cache_seqlens.to(q.device)
if k is not None:
sn = k.shape[1]
pos = cache_seqlens.view(b, 1).long() + torch.arange(sn, device=q.device).view(1, sn)
bidx = torch.arange(b, device=q.device).view(b, 1).expand(b, sn)
k_cache[bidx, pos] = k
v_cache[bidx, pos] = v
total = cache_seqlens.long() + sn
else:
total = cache_seqlens.long()
max_t = int(total.max())
kk = k_cache[:, :max_t]
vv = v_cache[:, :max_t]
if hq != hk:
rep = hq // hk
kk = kk.repeat_interleave(rep, dim=2)
vv = vv.repeat_interleave(rep, dim=2)
qt = q.transpose(1, 2) # (b, hq, sq, d)
kt = kk.transpose(1, 2) # (b, hq, T, d)
vt = vv.transpose(1, 2)
t_idx = torch.arange(max_t, device=q.device).view(1, 1, 1, max_t)
# query j sits at absolute position total-sq+j; causal => attend to <= that
limit = (total.view(b, 1, 1, 1) - sq
+ torch.arange(sq, device=q.device).view(1, 1, sq, 1))
mask = t_idx <= limit # (b, 1, sq, T) broadcast over heads
y = torch.nn.functional.scaled_dot_product_attention(
qt, kt, vt, attn_mask=mask)
return y.transpose(1, 2).contiguous()'''
assert anchor in src, "anchor not found - upstream changed"
src = src.replace(anchor, replacement, 1)
env_anchor = 'FISH_BATCH_INVARIANT = os.getenv("FISH_BATCH_INVARIANT", "false").lower() in ('
env_add = ('FISH_FORCE_SDPA = os.getenv("FISH_FORCE_SDPA", "false").lower() in '
'("true", "1", "yes")\n')
assert env_anchor in src
src = src.replace(env_anchor, env_add + env_anchor, 1)
P.write_text(src)
print("audio_decoder patched (FISH_FORCE_SDPA gate)")
import ast
ast.parse(src)
print("syntax OK")