@@ -280,39 +280,41 @@ def compute_layer_kv_cache_shape_bytes(
280280 return (num_blocks , spec .num_heads , num_states , spec .state_content_size_bytes )
281281
282282
283- def layer_kv_cache_strides (
283+ def compute_layout_strides (
284284 spec : KVCacheSpec ,
285285 num_blocks : int ,
286286 num_layers : int ,
287287 layout : KVCacheLayout ,
288288 block_size : int | None = None ,
289- ) -> tuple [ int , int ]:
290- """Byte ``(layer_stride, block_stride)`` of a dense ``[L, B, H, N, C]``
291- allocation in ``layout`` order."""
289+ packed_block_stride : int | None = None ,
290+ ) -> tuple [ int , ...]:
291+ """Byte strides in logical ``[L, B, H, N, C]`` axis order."""
292292 shape = (
293293 num_layers ,
294294 * compute_layer_kv_cache_shape_bytes (spec , num_blocks , block_size ),
295295 )
296296 stride_order = layout .stride_order
297297 physical_shape = tuple (shape [i ] for i in stride_order )
298- dense = torch .empty (physical_shape , device = "meta" ).stride ()
299298 inv_order = [stride_order .index (i ) for i in range (5 )]
300- layer_stride = dense [inv_order [_DIM_L ]]
301- block_stride = dense [inv_order [_DIM_B ]]
302299
303- if padded := getattr (spec , "page_size_padded" , None ):
300+ padded = getattr (spec , "page_size_padded" , None )
301+ if padded is not None :
304302 assert block_size is None or block_size == spec .block_size , (
305303 "Padded KV pages do not support kernel block splitting."
306304 )
307305 assert {inv_order [_DIM_L ], inv_order [_DIM_B ]} == {0 , 1 }, (
308306 f"Padded KV pages need L and B outermost, got { layout .name } ."
309307 )
310- # Padding widens every page, so the strides that step over whole
311- # pages scale with it.
312- page = prod (shape [2 :])
313- layer_stride = layer_stride // page * padded
314- block_stride = block_stride // page * padded
315- return layer_stride , block_stride
308+
309+ logical_tail = prod (physical_shape [2 :])
310+ storage_tail = padded if padded is not None else logical_tail
311+ physical = torch .empty ((* physical_shape [:2 ], storage_tail ), device = "meta" )
312+ strides = list (
313+ physical [..., :logical_tail ].view (physical_shape ).permute (* inv_order ).stride ()
314+ )
315+ if packed_block_stride is not None and inv_order [_DIM_B ] < inv_order [_DIM_L ]:
316+ strides [_DIM_B ] = packed_block_stride
317+ return tuple (strides )
316318
317319
318320def reshape_kv_cache (
@@ -338,21 +340,21 @@ def reshape_kv_cache(
338340 # e.g. BHLNC's head stride spans the layers it interleaves); the caller's
339341 # strides place the layers and blocks themselves.
340342 logical_shape = (num_layers , * shape_bytes )
341- stride_order = layout .stride_order
342- physical_shape = tuple (logical_shape [i ] for i in stride_order )
343- inv_order = [stride_order .index (i ) for i in range (5 )]
344- strides = list (torch .empty (physical_shape , device = "meta" ).stride ())
345- strides [inv_order [_DIM_L ]] = layer_stride
346- strides [inv_order [_DIM_B ]] = block_stride
343+ strides = list (
344+ compute_layout_strides (
345+ spec , num_blocks , num_layers , layout , block_size = block_size
346+ )
347+ )
348+ strides [_DIM_L ] = layer_stride
349+ strides [_DIM_B ] = block_stride
347350 dtype = getattr (spec , "dtype" , None )
348351
349- cache = torch .as_strided (
352+ cache_logical_5d = torch .as_strided (
350353 raw ,
351- size = physical_shape ,
354+ size = logical_shape ,
352355 stride = tuple (strides ),
353356 storage_offset = raw .storage_offset () + offset ,
354357 )
355- cache_logical_5d = cache .permute (* inv_order )
356358
357359 views = []
358360 for layer_idx in range (num_layers ):
0 commit comments