25 lines
1.0 KiB
Python
25 lines
1.0 KiB
Python
path = '/usr/local/lib/python3.12/dist-packages/vllm/v1/attention/backends/triton_attn.py'
|
|
with open(path, 'r') as f:
|
|
content = f.read()
|
|
|
|
old = ''' def validate_head_size(cls, head_size: int) -> None:
|
|
# Triton Attention supports any head size above 32
|
|
if head_size < 32:
|
|
raise ValueError(
|
|
f"Head size {head_size} is not supported by TritonAttention."
|
|
f"Head sizes need to be larger or equal 32 for this backend. "
|
|
"Set VLLM_ATTENTION_BACKEND=FLEX_ATTENTION to use "
|
|
"FlexAttention backend which supports all head sizes.")'''
|
|
|
|
new = ''' def validate_head_size(cls, head_size: int) -> None:
|
|
# PATCH: allow all head sizes (Triton compiles at runtime)
|
|
return'''
|
|
|
|
if old in content:
|
|
content = content.replace(old, new)
|
|
with open(path, 'w') as f:
|
|
f.write(content)
|
|
print('patch_triton: validate_head_size bypassed successfully')
|
|
else:
|
|
print('patch_triton: WARNING - pattern not found, patch skipped')
|