use tract_core::internal::*;
use tract_core::ops::konst::Const;
use tract_core::tract_linalg::block_quant::*;
use tract_gpu::fact::DeviceFact;
use tract_gpu::rule_ensure;
use tract_gpu::tensor::{DeviceTensor, DeviceTensorExt, IntoDevice, OwnedDeviceTensor};
use tract_gpu::utils::as_q40_tensor;
use crate::Q40_ROW_PADDING;
use crate::ops::{CudaFusedAxisOp, CudaGgmlGemm};
use crate::tensor::CudaTensor;
use crate::utils::pad_q40;
use tract_gpu::ops::change_axes::GpuAxisOp;
fn effective_gemm_shape(
model: &TypedModel,
node: &TypedNode,
shape: &[usize],
) -> TractResult<Option<TVec<usize>>> {
let mut cursor = node;
let mut effective_shape: TVec<usize> = shape.into();
while let Some(succ) = model.single_succ(cursor.id)? {
if succ.op_is::<CudaGgmlGemm>() {
return Ok(Some(effective_shape));
}
if let Some(fao) = succ.op_as::<CudaFusedAxisOp>() {
if fao.op.is::<CudaGgmlGemm>() {
let weight_inlet = succ.inputs.iter().position(|i| i.node == cursor.id).unwrap();
for axis_op in &fao.grouped_axis_ops[weight_inlet] {
axis_op.inner.change_shape_array(&mut effective_shape, false)?;
}
return Ok(Some(effective_shape));
}
}
if let Some(axis_op) = succ.op_as::<GpuAxisOp>() {
axis_op.inner.change_shape_array(&mut effective_shape, false)?;
cursor = succ;
continue;
}
break;
}
Ok(None)
}
pub fn pad_q40_weights(
_ctx: &(),
model: &TypedModel,
node: &TypedNode,
_node_name: &str,
op: &Const,
) -> TractResult<Option<TypedModelPatch>> {
let Some(dev_tensor) = op.val().to_device_tensor().ok() else {
return Ok(None);
};
let DeviceTensor::Owned(t) = dev_tensor else {
return Ok(None);
};
let Some(cuda_tensor) = t.downcast_ref::<CudaTensor>() else {
return Ok(None);
};
let bqf = cuda_tensor
.exotic_fact()
.and_then(|of| of.downcast_ref::<BlockQuantFact>())
.filter(|bqf| bqf.format.dyn_eq(&Q4_0));
rule_ensure!(bqf.is_some());
let bqf = bqf.unwrap();
let Some(effective_shape) = effective_gemm_shape(model, node, bqf.shape())? else {
return Ok(None);
};
let effective_k = *effective_shape.last().unwrap();
rule_ensure!(effective_k % Q40_ROW_PADDING != 0);
let host_tensor = dev_tensor.to_host()?.into_tensor();
let bqs = as_q40_tensor(&host_tensor).expect("expected Q4_0 tensor view");
let total_elements: usize = bqf.shape().iter().product();
let flat_m = total_elements / effective_k;
let padded_bqs = pad_q40(bqs, flat_m, effective_k)?;
let padded_k = effective_k.next_multiple_of(Q40_ROW_PADDING);
let mut padded_shape = effective_shape.clone();
*padded_shape.last_mut().unwrap() = padded_k;
let padded_bqf = BlockQuantFact::new(
tract_core::dyn_clone::clone_box(padded_bqs.format()),
padded_shape.clone(),
);
let padded_fact =
TypedFact::dt_shape(f32::datum_type(), &padded_shape).with_exotic_fact(padded_bqf);
let padded_tensor =
padded_bqs.into_tensor_with_shape(f32::datum_type(), &padded_shape).into_arc_tensor();
let new_const = Const::new_with_exotic_fact(
padded_tensor.into_device()?.into_tensor().into_arc_tensor(),
Box::new(DeviceFact::from_host(padded_fact)?),
)?;
let mut patch = TypedModelPatch::default();
let wire = patch.wire_node(&node.name, new_const, &[])?[0];
let mut cursor = node;
let mut obliterate_ids: TVec<usize> = tvec![node.id];
while let Some(succ) = model.single_succ(cursor.id)? {
if succ.op_is::<GpuAxisOp>() {
obliterate_ids.push(succ.id);
cursor = succ;
continue;
}
break;
}
let gemm_succ = model.single_succ(cursor.id)?.unwrap();
let weight_outlet: OutletId = cursor.id.into();
let weight_inlet = gemm_succ
.inputs
.iter()
.position(|i| i.node == weight_outlet.node && i.slot == weight_outlet.slot)
.unwrap();
let mut gemm_inputs: TVec<OutletId> = tvec![];
for (ix, input) in gemm_succ.inputs.iter().enumerate() {
if ix == weight_inlet {
gemm_inputs.push(wire);
} else {
gemm_inputs.push(patch.tap_model(model, *input)?);
}
}
let gemm_op: Box<dyn TypedOp> = if let Some(fao) = gemm_succ.op_as::<CudaFusedAxisOp>() {
let mut axis_ops = fao.grouped_axis_ops.clone();
axis_ops[weight_inlet].clear();
Box::new(CudaFusedAxisOp::new(axis_ops, fao.op.clone()))
} else {
gemm_succ.op.clone()
};
let gemm_out = patch.wire_node(&gemm_succ.name, gemm_op, &gemm_inputs)?;
patch.shunt_outside(model, gemm_succ.id.into(), gemm_out[0])?;
obliterate_ids.push(gemm_succ.id);
for id in obliterate_ids {
patch.obliterate(id)?;
}
Ok(Some(patch))
}