use std::any::TypeId;
use std::collections::HashMap;
use std::sync::OnceLock;
use crate::context::cuda_context;
use crate::ops::wire_cuda_conv;
use crate::{kernels, ops, rewrite_rules};
use DatumType::{F16, F32};
use tract_core::dyn_clone::clone_box;
use tract_core::internal::*;
use tract_core::model::translator::Translate;
use tract_core::ops::cnn::conv::rewrite_kernel_conv_in_oihw;
use tract_core::ops::cnn::{Conv, rewrite_conv_with_n_axis};
use tract_core::ops::einsum::prefix_matmul::{PrefixMatMul, rewrite_einsum_to_prefix_matmul};
use tract_core::ops::konst::Const;
tract_core::declare_knob!(
TRACT_CUDA_FORCE_CPU,
String,
String::new(),
"Comma-separated node-name substrings forced onto the CPU instead of CUDA."
);
tract_core::declare_knob!(
TRACT_CUDA_TRANSLATE_DEBUG,
bool,
false,
"Log nodes that fail output-fact validation during CUDA translation."
);
use tract_core::ops::nn::{Reduce, rewrite_nearest_upsample_to_broadcast};
use tract_core::tract_linalg::block_quant::Q4_0;
use tract_core::transform::ModelTransform;
use tract_gpu::fact::{DeviceFact, DeviceTypedFactExt};
use tract_gpu::rewrite_rules::rewire_syncs::rewire_syncs;
use tract_gpu::rewrite_rules::rms_norm::{
fuse_rms_norm_residual, fuse_rms_norm_scale, fuse_rms_norm_split_scale,
fuse_scaled_rms_norm_in_cast, fuse_scaled_rms_norm_out_cast, remove_rms_norm_cast,
};
use tract_gpu::sync::{DeviceSyncKind, sync_inputs_if_required, sync_model_outputs_if_required};
use tract_gpu::tensor::{DeviceTensor, IntoDevice};
use tract_gpu::utils::as_quant_fact;
use tract_transformers::ops::sdpa::Sdpa;
pub struct CudaOpTranslator {
pub type_id: TypeId,
#[allow(clippy::type_complexity)]
pub try_make: fn(&TypedModel, &TypedNode) -> TractResult<Option<Box<dyn TypedOp>>>,
}
inventory::collect!(CudaOpTranslator);
#[macro_export]
macro_rules! register_cuda_op {
($op_type:ty, |$source:ident, $node:ident, $op:ident| $body:expr) => {
inventory::submit! {
$crate::transform::CudaOpTranslator {
type_id: std::any::TypeId::of::<$op_type>(),
try_make: |$source, $node| {
let Some($op) = $node.op_as::<$op_type>() else {
return Ok(None);
};
$body
},
}
}
};
}
#[derive(Debug, Default)]
pub struct CudaTransform;
impl ModelTransform for CudaTransform {
fn name(&self) -> StaticName {
"cuda-transform".into()
}
fn transform(&self, model: &mut TypedModel) -> TractResult<()> {
self.transform_up_to_phase(model, usize::MAX)
}
}
impl CudaTransform {
pub fn transform_up_to_phase(
&self,
model: &mut TypedModel,
stop_at_phase: usize,
) -> TractResult<()> {
cuda_context();
rewire_sdpa_cuda(model)?;
rewrite_einsum_to_prefix_matmul(model, false)?;
if stop_at_phase == 0 {
return Ok(());
}
Rewriter::default()
.with_rule_for("untranspose_matmul_output", rewrite_rules::untranspose_matmul_output)
.with_rule_for("add_broadcast_pre_matmul", rewrite_rules::add_broadcast_pre_matmul)
.with_rule_for("rewrite_kernel_conv_in_oihw", rewrite_kernel_conv_in_oihw)
.with_rule_for("rewrite_conv_with_n_axis", rewrite_conv_with_n_axis)
.with_rule_for("nearest_upsample_to_broadcast", rewrite_nearest_upsample_to_broadcast)
.rewrite(&(), model)?;
Rewriter::default()
.with_rule_for("fuse_rms_norm_scale", fuse_rms_norm_scale)
.with_rule_for("fuse_rms_norm_split_scale", fuse_rms_norm_split_scale)
.with_rule_for("remove_rms_norm_cast", remove_rms_norm_cast)
.with_rule_for("fuse_scaled_rms_norm_in_cast", fuse_scaled_rms_norm_in_cast)
.with_rule_for("fuse_scaled_rms_norm_out_cast", fuse_scaled_rms_norm_out_cast)
.with_rule_for("split_multi_axis_reduce", split_multi_axis_reduce)
.rewrite(&(), model)?;
if stop_at_phase == 1 {
return Ok(());
}
*model = self.translate_model(model)?;
if stop_at_phase == 2 {
return Ok(());
}
Rewriter::default()
.with_rule_for("fuse_rms_norm_residual", fuse_rms_norm_residual)
.rewrite(&(), model)?;
Rewriter::default()
.with_rule_for("fuse_move_axis", rewrite_rules::fuse_move_axis)
.rewrite(&(), model)?;
Rewriter::default()
.with_rule_for("fuse_output_axis_op", tract_gpu::ops::fused_output::fuse_output_axis_op)
.rewrite(&(), model)?;
Rewriter::default()
.with_rule_for("fuse_window_into_gemm", rewrite_rules::fuse_window_into_gemm)
.rewrite(&(), model)?;
Rewriter::default()
.with_rule_for("fuse_axis_op", rewrite_rules::fuse_axis_op)
.rewrite(&(), model)?;
rewire_syncs(model)?;
Rewriter::default()
.with_rule_for("pad_q40_weights", rewrite_rules::pad_q40_weights)
.rewrite(&(), model)?;
Ok(())
}
}
fn try_make_cuda_op(
source: &TypedModel,
node: &TypedNode,
) -> TractResult<Option<Box<dyn TypedOp>>> {
type TranslateFn = fn(&TypedModel, &TypedNode) -> TractResult<Option<Box<dyn TypedOp>>>;
static MAP: OnceLock<HashMap<TypeId, Vec<TranslateFn>>> = OnceLock::new();
let map = MAP.get_or_init(|| {
let mut m: HashMap<TypeId, Vec<TranslateFn>> = HashMap::new();
for t in inventory::iter::<CudaOpTranslator> {
m.entry(t.type_id).or_default().push(t.try_make);
}
m
});
let input_facts = source.node_input_facts(node.id)?;
rule_if!(input_facts.iter().all(|f| DeviceTensor::is_supported_dt(f.datum_type)));
if let Some(op) = tract_gpu::ops::copy_based::try_make_copy_based_op(source, node)? {
return Ok(Some(op));
}
if let Some(fns) = map.get(&(*node.op).type_id()) {
for f in fns {
if let Some(op) = f(source, node)? {
return Ok(Some(op));
}
}
}
Ok(None)
}
fn convert_const(op: &Const) -> TractResult<Const> {
let typed_fact: TypedFact = Arc::clone(op.val()).try_into()?;
let cuda_fact = if let Some(of) = op.exotic_fact() {
DeviceFact::from_host(typed_fact.with_exotic_fact(clone_box(of)))?
} else {
DeviceFact::from_host(typed_fact)?
};
let cuda_const = op.val().clone().into_device()?.into_tensor().into_arc_tensor();
Const::new_with_exotic_fact(cuda_const, Box::new(cuda_fact))
}
pub(crate) fn cuda_cast_new(to: DatumType) -> Option<tract_gpu::ops::cast::GpuCast> {
tract_gpu::ops::cast::GpuCast::new(
to,
"Cuda",
kernels::array::cuda_cast_dispatch,
kernels::array::Cast::is_supported_dt,
)
}
fn can_convert_to_cuda_gemm(facts: &[TypedFact]) -> bool {
assert!(facts.len() == 2, "Ggml: Expected 2 inputs for Matmul");
let regular_types_support = facts[0].is_plain()
&& facts[1].is_plain()
&& matches!(
(facts[0].datum_type, facts[1].datum_type),
(F32, F32) | (F16, F16) | (F16, F32)
);
regular_types_support
|| (as_quant_fact(&facts[1], &Q4_0).is_some() && matches!(facts[0].datum_type, F16 | F32))
}
fn convert_matmul_to_cuda(
model: &TypedModel,
node: &TypedNode,
target: &mut TypedModel,
inputs: &mut [OutletId],
op: &PrefixMatMul,
) -> TractResult<TVec<OutletId>> {
let mut input_facts = model.node_input_facts(node.id)?;
let mut swap_inputs = false;
if !can_convert_to_cuda_gemm(&[input_facts[0].clone(), input_facts[1].clone()])
&& can_convert_to_cuda_gemm(&[input_facts[1].clone(), input_facts[0].clone()])
{
input_facts.swap(0, 1);
inputs.swap(0, 1);
swap_inputs = true;
}
let act_fact = input_facts[0];
let weight_fact = input_facts[1];
let outlets = inputs.split_at_mut(1);
let act_outlet = &mut outlets.0[0];
let weights_outlet = &mut outlets.1[0];
let transpose_act = if swap_inputs { !op.transpose_b } else { op.transpose_a };
let transpose_weight = if swap_inputs { !op.transpose_a } else { op.transpose_b };
if transpose_act {
let rank = act_fact.rank();
let perm_act_op =
tract_gpu::ops::change_axes::GpuAxisOp::new(AxisOp::Move(rank - 2, rank - 1));
let perm_act_name = node.name.clone() + ".perm_activs";
*act_outlet = target.wire_node(perm_act_name, perm_act_op, &[*act_outlet])?[0];
}
if act_fact.datum_type == DatumType::F16 && as_quant_fact(weight_fact, &Q4_0).is_some() {
let in_cast_op = cuda_cast_new(DatumType::F32).unwrap();
*act_outlet =
target.wire_node(node.name.clone() + ".in_cast", in_cast_op, &[*act_outlet])?[0];
} else if act_fact.datum_type == DatumType::F16 && weight_fact.datum_type == DatumType::F32 {
let in_cast_op = cuda_cast_new(DatumType::F16).unwrap();
*weights_outlet =
target.wire_node(node.name.clone() + ".in_cast", in_cast_op, &[*weights_outlet])?[0];
}
if !transpose_weight {
ensure!(as_quant_fact(weight_fact, &Q4_0).is_none(), "Cannot transpose Q40 tensor");
let rank = weight_fact.rank();
let perm_weights_op =
tract_gpu::ops::change_axes::GpuAxisOp::new(AxisOp::Move(rank - 2, rank - 1));
let perm_weights_name = node.name.clone() + ".perm_weights";
*weights_outlet =
target.wire_node(perm_weights_name, perm_weights_op, &[*weights_outlet])?[0];
}
if as_quant_fact(weight_fact, &Q4_0).is_some() {
let device_fact = target.outlet_fact(*act_outlet)?.to_device_fact()?;
let quant_op = ops::CudaGgmlQuantQ81::new(device_fact.shape.clone())?;
*act_outlet =
target.wire_node(node.name.clone() + ".quant_activs", quant_op, &[*act_outlet])?[0];
}
let mut matmul_output =
target.wire_node(node.name.clone(), *Box::new(ops::CudaGgmlGemm), inputs)?;
if swap_inputs {
let out_fact = target.outlet_fact(matmul_output[0])?;
let rank = &out_fact
.exotic_fact
.clone()
.map(|fact| fact.clarify_dt_shape().unwrap().1.len())
.unwrap();
let perm_out_op =
tract_gpu::ops::change_axes::GpuAxisOp::new(AxisOp::Move(rank - 2, rank - 1));
matmul_output =
target.wire_node(node.name.clone() + ".perm_out", perm_out_op, &matmul_output)?;
}
let out_fact = target.outlet_fact(matmul_output[0])?;
let out_dt = out_fact.as_device_fact().map(|f| f.datum_type).unwrap_or(out_fact.datum_type);
let expected_dt = model.node_output_facts(node.id)?[0].datum_type;
if out_dt != expected_dt {
ensure!(
kernels::array::Cast::is_supported_dt(out_dt),
"Matmul output type cannot be casted to expected type"
);
let cast_op = cuda_cast_new(model.node_output_facts(node.id)?[0].datum_type).unwrap();
matmul_output =
target.wire_node(node.name.clone() + ".out_cast", cast_op, &matmul_output)?
}
Ok(matmul_output)
}
fn cuda_flash_attn_supported(facts: &[&TypedFact]) -> bool {
if !matches!(facts.len(), 3 | 4) {
return false;
}
let [qf, kf, vf] = [facts[0], facts[1], facts[2]];
if kf.datum_type() != vf.datum_type() || kf.shape != vf.shape {
return false;
}
if [qf, kf, vf].iter().any(|f| !matches!(f.rank(), 3 | 4)) {
return false;
}
kf.shape[kf.rank() - 1].to_i64().is_ok_and(|head_dim| matches!(head_dim, 64 | 128))
}
fn flatten_unfused_sdpa(
_ctx: &(),
model: &TypedModel,
node: &TypedNode,
_name: &str,
op: &Sdpa,
) -> TractResult<Option<TypedModelPatch>> {
if cuda_flash_attn_supported(&model.node_input_facts(node.id)?) {
Ok(None)
} else {
op.patch_sdpa(model, node)
}
}
fn rewire_sdpa_cuda(model: &mut TypedModel) -> TractResult<()> {
Rewriter::default()
.with_rule_for("flatten-unfused-sdpa", flatten_unfused_sdpa)
.rewrite(&(), model)
}
fn convert_sdpa_to_cuda_flash_attn(
model: &TypedModel,
node: &TypedNode,
target: &mut TypedModel,
inputs: &mut [OutletId],
op: &Sdpa,
) -> TractResult<TVec<OutletId>> {
let facts = model.node_input_facts(node.id)?;
ensure!(cuda_flash_attn_supported(&facts), "Sdpa shape the flash-attention kernel cannot take");
let [qf, kf, vf] = [facts[0], facts[1], facts[2]];
let mask_fact = if facts.len() == 4 { Some(facts[3]) } else { None };
let (q, k, v, m_opt) = match &mut inputs[..] {
[q, k, v, m, ..] => (q, k, v, Some(m)),
[q, k, v] => (q, k, v, None),
_ => bail!("unexpected number of inputs"),
};
fn name(base: &str, suffix: &str) -> String {
format!("{base}{suffix}")
}
fn mut_cast(
target: &mut TypedModel,
node_name: &str,
dst: &mut OutletId,
have: DatumType,
want: DatumType,
suffix: &str,
) -> TractResult<()> {
if have != want {
*dst =
target.wire_node(name(node_name, suffix), cuda_cast_new(want).unwrap(), &[*dst])?
[0];
}
Ok(())
}
fn add_head_axis_if_rank3(
target: &mut TypedModel,
node_name: &str,
dst: &mut OutletId,
fact: &TypedFact,
suffix: &str,
) -> TractResult<bool> {
if fact.rank() == 3 {
let ax = tract_gpu::ops::change_axes::GpuAxisOp::new(AxisOp::Add(1));
*dst = target.wire_node(name(node_name, suffix), ax, &[*dst])?[0];
Ok(true)
} else {
ensure!(fact.rank() == 4, "Q/K/V must be rank 3 or 4");
Ok(false)
}
}
let q_dt = qf.datum_type().unwrap();
let kv_dt = kf.datum_type().unwrap();
mut_cast(target, &node.name, k, kv_dt, DatumType::F16, ".cast_k")?;
mut_cast(target, &node.name, v, kv_dt, DatumType::F16, ".cast_v")?;
mut_cast(target, &node.name, q, q_dt, DatumType::F16, ".cast_q")?;
let mut added_head_axis = false;
added_head_axis |= add_head_axis_if_rank3(target, &node.name, q, qf, ".reshape_q")?;
added_head_axis |= add_head_axis_if_rank3(target, &node.name, k, kf, ".reshape_k")?;
added_head_axis |= add_head_axis_if_rank3(target, &node.name, v, vf, ".reshape_v")?;
let out_dim = kf.shape[kf.rank() - 1].to_i64()?;
if let Some(mf) = mask_fact {
let m = m_opt.unwrap();
mut_cast(target, &node.name, m, mf.datum_type().unwrap(), DatumType::F16, ".cast_m")?;
if mf.rank() != 4 {
let ax = tract_gpu::ops::change_axes::GpuAxisOp::new(AxisOp::Add(1));
*m = target.wire_node(name(&node.name, ".reshape_m"), ax, &[*m])?[0];
}
}
let scale = op
.scale
.as_ref()
.map(|s| *s.try_as_plain_ram().unwrap().to_scalar::<f32>().unwrap())
.unwrap_or(1.0 / (out_dim as f32).sqrt());
let sdpa = ops::CudaFlashAttention::new(scale, op.is_causal);
let mut out = target.wire_node(node.name.clone(), sdpa, inputs)?;
if added_head_axis {
out = target.wire_node(
name(&node.name, ".reshape_out"),
tract_gpu::ops::change_axes::GpuAxisOp::new(AxisOp::Rm(1)),
&out,
)?;
}
if q_dt != DatumType::F16 {
out =
target.wire_node(name(&node.name, ".cast_out"), cuda_cast_new(q_dt).unwrap(), &out)?;
}
Ok(out)
}
impl Translate<TypedFact, Box<dyn TypedOp>, TypedFact, Box<dyn TypedOp>> for CudaTransform {
fn translate_node(
&self,
source: &TypedModel,
node: &TypedNode,
target: &mut TypedModel,
mapping: &HashMap<OutletId, OutletId>,
) -> TractResult<TVec<OutletId>> {
let input_facts = source.node_input_facts(node.id)?;
if let Some(op) = node.op_as::<PrefixMatMul>() {
let facts: Vec<TypedFact> = input_facts.iter().map(|f| (*f).clone()).collect();
if !op.transpose_c
&& op.quantize_output.is_none()
&& (can_convert_to_cuda_gemm(&facts)
|| can_convert_to_cuda_gemm(&[facts[1].clone(), facts[0].clone()]))
{
let mut device_inputs =
sync_inputs_if_required(target, node, mapping, DeviceSyncKind::ToDevice)?;
let outlet_ids =
convert_matmul_to_cuda(source, node, target, &mut device_inputs, op)?;
return sync_model_outputs_if_required(source, node, target, outlet_ids);
}
}
if let Some(op) = node.op_as::<Sdpa>()
&& cuda_flash_attn_supported(&input_facts)
{
let mut device_inputs =
sync_inputs_if_required(target, node, mapping, DeviceSyncKind::ToDevice)?;
let outlet_ids =
convert_sdpa_to_cuda_flash_attn(source, node, target, &mut device_inputs, op)?;
return sync_model_outputs_if_required(source, node, target, outlet_ids);
}
if let Some(conv) = node.op_as::<Conv>()
&& input_facts.iter().all(|f| DeviceTensor::is_supported_dt(f.datum_type))
&& matches!(input_facts[0].datum_type, F16 | F32)
{
let device_inputs =
sync_inputs_if_required(target, node, mapping, DeviceSyncKind::ToDevice)?;
let outlet_ids = wire_cuda_conv(source, node, target, &device_inputs, conv)?;
return sync_model_outputs_if_required(source, node, target, outlet_ids);
}
if let Some(op) = node.op_as::<Const>()
&& DeviceTensor::is_supported_dt(op.val().datum_type())
{
let device_inputs =
sync_inputs_if_required(target, node, mapping, DeviceSyncKind::ToDevice)?;
let outlet_ids =
target.wire_node(node.name.clone(), convert_const(op)?, &device_inputs)?;
return sync_model_outputs_if_required(source, node, target, outlet_ids);
}
let target_inputs: TVec<TypedFact> = node
.inputs
.iter()
.map(|i| target.outlet_fact(mapping[i]).cloned())
.collect::<TractResult<_>>()?;
let target_inputs_post_sync: TVec<TypedFact> = target_inputs
.iter()
.map(|f| -> TractResult<TypedFact> {
if f.as_device_fact().is_some() {
Ok(f.clone())
} else {
Ok(tract_gpu::fact::DeviceFact::from_host(f.clone())?.into_exotic_fact())
}
})
.collect::<TractResult<_>>()?;
let target_input_post_sync_refs: TVec<&TypedFact> =
target_inputs_post_sync.iter().collect();
let force_cpu = TRACT_CUDA_FORCE_CPU
.get()
.split(',')
.any(|pat| !pat.is_empty() && node.name.contains(pat));
let maybe_gpu_op = if force_cpu { None } else { try_make_cuda_op(source, node)? };
if let Some(ref op) = maybe_gpu_op
&& TRACT_CUDA_TRANSLATE_DEBUG.get()
&& let Err(e) = op.output_facts(&target_input_post_sync_refs)
{
eprintln!(
"cuda-translate-fallback: {} ({}) inputs={:?} -> {e:?}",
node.name,
op.name(),
target_inputs_post_sync,
);
}
if let Some(gpu_op) = maybe_gpu_op
&& gpu_op.output_facts(&target_input_post_sync_refs).is_ok()
{
let device_inputs =
sync_inputs_if_required(target, node, mapping, DeviceSyncKind::ToDevice)?;
let outlet_ids = target.wire_node(node.name.clone(), gpu_op, &device_inputs)?;
sync_model_outputs_if_required(source, node, target, outlet_ids)
} else {
let cpu_inputs =
sync_inputs_if_required(target, node, mapping, DeviceSyncKind::ToHost)?;
target.wire_node(&node.name, node.op.clone(), &cpu_inputs)
}
}
}
fn split_multi_axis_reduce(
_ctx: &(),
model: &TypedModel,
node: &TypedNode,
node_name: &str,
op: &Reduce,
) -> TractResult<Option<TypedModelPatch>> {
rule_if!(op.axes.len() > 1);
use tract_core::ops::nn::Reducer::*;
rule_if!(matches!(op.reducer, Sum | Prod | Min | Max | Any | All));
let mut patch = TypedModelPatch::default();
let mut wire = patch.tap_model(model, node.inputs[0])?;
let mut axes = op.axes.clone();
axes.sort();
for (i, &axis) in axes.iter().rev().enumerate() {
let single = Reduce { axes: tvec![axis], reducer: op.reducer };
wire = patch.wire_node(format!("{node_name}.axis_{i}"), single, &[wire])?[0];
}
patch.shunt_outside(model, node.id.into(), wire)?;
Ok(Some(patch))
}
#[cfg(test)]
mod test {
use super::*;
#[test]
fn test_prefix_matmul_transform_f32_f16() -> TractResult<()> {
let mut model = TypedModel::default();
let (b, m, k, n) = (1, 16, 128, 32);
let a_fact = TypedFact::dt_shape(DatumType::F32, [b, m, k]);
let b_fact = TypedFact::dt_shape(DatumType::F16, [b, k, n]);
let source_a = model.add_source("a", a_fact)?;
let source_b = model.add_source("b", b_fact)?;
let op = PrefixMatMul {
transpose_a: false,
transpose_b: false,
transpose_c: false,
quantize_output: None,
operating_dt: Some(DatumType::F32),
};
let matmul_out = model.wire_node("matmul", op, &[source_a, source_b])?;
model.select_output_outlets(&matmul_out)?;
let tensor_a = Tensor::zero::<f32>(&[b, m, k])?;
let tensor_b = Tensor::zero::<f16>(&[b, k, n])?;
let inputs = tvec!(tensor_a.into(), tensor_b.into());
let transform = CudaTransform;
transform.transform(&mut model)?;
let cuda_runnable = model.into_runnable()?;
let _ = cuda_runnable.run(inputs)?;
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
fn sdpa_model(head_dim: usize) -> TractResult<(TypedModel, TVec<TValue>)> {
let (heads, seq) = (4usize, 32usize);
let shape = [1, heads, seq, head_dim];
let fact = f32::fact(shape.iter().map(|d| TDim::from(*d as i64)).collect::<TVec<_>>());
let mut model = TypedModel::default();
let q = model.add_source("q", fact.clone())?;
let k = model.add_source("k", fact.clone())?;
let v = model.add_source("v", fact.clone())?;
let out = model.wire_node(
"sdpa",
Sdpa {
scale: None,
datum_type: f32::datum_type(),
acc_datum_type: f32::datum_type(),
is_causal: false,
},
&[q, k, v],
)?;
model.select_output_outlets(&out)?;
let mk = |seed: i64| -> TractResult<TValue> {
let data: Vec<f32> = (0..heads * seq * head_dim)
.map(|i| (((i as i64 * 2654435761 + seed).rem_euclid(1000)) as f32 / 1000.0) - 0.5)
.collect();
Ok(Tensor::from_shape(&shape, &data)?.into_tvalue())
};
Ok((model, tvec!(mk(1)?, mk(2)?, mk(3)?)))
}
#[test]
fn sdpa_at_supported_head_dim_routes_to_flash_attention() -> TractResult<()> {
let (model, _) = sdpa_model(64)?;
let cuda = CudaTransform.transform_into(model)?;
assert!(
cuda.nodes().iter().any(|n| n.op_is::<ops::CudaFlashAttention>()),
"a head dim of 64 should reach CudaFlashAttention"
);
Ok(())
}
#[test]
fn sdpa_at_unsupported_head_dim_explodes_and_matches_cpu() -> TractResult<()> {
let (model, inputs) = sdpa_model(8)?;
let expected = model.clone().into_runnable()?.run(inputs.clone())?;
let cuda = CudaTransform.transform_into(model)?;
assert!(
!cuda.nodes().iter().any(|n| n.op_is::<ops::CudaFlashAttention>()),
"a head dim of 8 must not reach CudaFlashAttention"
);
let found = cuda.into_runnable()?.run(inputs)?;
found[0].close_enough(&expected[0], Approximation::Approximate)?;
Ok(())
}
}