Download scripts/patch_audio_decoder.py from ehabnegm/s2pro-egy: direct link, hf CLI and curl.
- Browser
- Download file 3.46 kB
-
https://huggingface.co/ehabnegm/s2pro-egy/resolve/main/scripts/patch_audio_decoder.py
- Command line
-
hf download hf://ehabnegm/s2pro-egy/scripts/patch_audio_decoder.py
-
curl -L -o patch_audio_decoder.py https://huggingface.co/ehabnegm/s2pro-egy/resolve/main/scripts/patch_audio_decoder.py
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") | |