[BASE] activation.py SiluAndMul: CCCL dispatch_transform CacheAsyncConfiguration output tensor caching
Source: cccl_upstream/cub/cub/device/dispatch/dispatch_transform.cuh Target: vllm/model_executor/layers/activation.py CCCL system design applied: - dispatch_transform.cuh CacheAsyncConfiguration: cache occupancy/config results across calls to avoid recomputation - Applied: cache output tensor when shape/dtype/device unchanged - BI-V100 has no async allocator → cudaMalloc is synchronous → caching avoids blocking the stream on every decode step - spread_out_items_per_thread: dynamic tile adjustment for occupancy → we only cache for stable decode shapes, not variable prefill
This commit is contained in:
@@ -33,7 +33,31 @@ class SiluAndMul(CustomOp):
|
||||
|
||||
d = x.shape[-1] // 2
|
||||
output_shape = (x.shape[:-1] + (d, ))
|
||||
out = torch.empty(output_shape, dtype=x.dtype, device=x.device)
|
||||
# ═══════════════════════════════════════════════════════════════
|
||||
# CCCL dispatch_transform.cuh CacheAsyncConfiguration pattern:
|
||||
# "This computation MUST NOT depend on any runtime state of the
|
||||
# current API invocation (like num_items), since the result
|
||||
# will be cached."
|
||||
#
|
||||
# For element-wise transforms, the output tensor shape is
|
||||
# deterministic from the input shape. During decode, input shape
|
||||
# is stable (num_seqs × hidden_dim doesn't change between steps).
|
||||
# Cache the output tensor to avoid cudaMalloc on every step.
|
||||
#
|
||||
# CCCL also uses spread_out_items_per_thread to dynamically
|
||||
# adjust tile size for small problems — analogously, we only
|
||||
# cache when shapes are stable (decode), not during prefill
|
||||
# where shapes vary per request.
|
||||
# ═══════════════════════════════════════════════════════════════
|
||||
_cache_key = (output_shape, x.dtype, x.device)
|
||||
_cached = getattr(self, '_out_cache', {}).get(_cache_key)
|
||||
if _cached is not None and _cached.shape == output_shape:
|
||||
out = _cached
|
||||
else:
|
||||
out = torch.empty(output_shape, dtype=x.dtype, device=x.device)
|
||||
if not hasattr(self, '_out_cache'):
|
||||
self._out_cache = {}
|
||||
self._out_cache[_cache_key] = out
|
||||
ops.silu_and_mul(out, x)
|
||||
return out
|
||||
|
||||
|
||||
Reference in New Issue
Block a user