shard_gdn_qkv_rows

Function shard_gdn_qkv_rows 

Source
pub fn shard_gdn_qkv_rows(
    src: DevicePtr,
    dims: &TpGdnDims,
    gpu: &dyn GpuBackend,
) -> Result<(DevicePtr, usize, usize)>
Expand description

Shard the [Q|K|V] (in_proj_qkv) BF16 weight [full_conv_dim, h] to the local rank’s [local_conv_dim, h], slicing Q, K and V independently by the local head range. Returns (ptr, local_rows, h).