shard_dense_weight

Function shard_dense_weight 

Source
pub fn shard_dense_weight(
    src: &DenseWeight,
    out_dim: usize,
    in_dim: usize,
    kind: TpShardKind,
    tp_rank: usize,
    tp_size: usize,
    gpu: &dyn GpuBackend,
) -> Result<(DenseWeight, usize, usize)>
Expand description

Convenience wrapper: shard a DenseWeight BF16 tensor. The source weight is freed by the caller — this fn allocates a new device buffer.