shard_gdn_qkvz_rows

Function shard_gdn_qkvz_rows 

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

Shard the concatenated [Q|K|V|Z] (in_proj_qkvz) BF16 weight [full_qkvz_out, h] to the local rank’s [local_qkvz_out, h], slicing all four segments independently. Returns (ptr, local_rows, h).