use crate::safe::{PhysicalDevice, SubgroupFeatureFlags};
use kiss_vulkan_vocab::{
Arith, ComponentType, CoopMatrix, CoopShape, OpClasses, Subgroup, VulkanTarget,
};
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct DeviceCapabilities {
pub default_subgroup: u32,
pub subgroup_range: Option<(u32, u32)>,
pub ops: OpClasses,
pub arith: Arith,
pub coop: Vec<CoopShape>,
}
impl DeviceCapabilities {
pub fn of(physical: &PhysicalDevice) -> Option<Self> {
let sg = physical.subgroup_properties()?;
let mut ops = OpClasses::NONE;
let supported = sg.supported_operations;
for (flag, class) in [
(SubgroupFeatureFlags::BASIC, OpClasses::BASIC),
(SubgroupFeatureFlags::VOTE, OpClasses::VOTE),
(SubgroupFeatureFlags::ARITHMETIC, OpClasses::ARITHMETIC),
(SubgroupFeatureFlags::BALLOT, OpClasses::BALLOT),
(SubgroupFeatureFlags::SHUFFLE, OpClasses::SHUFFLE),
(
SubgroupFeatureFlags::SHUFFLE_RELATIVE,
OpClasses::SHUFFLE_RELATIVE,
),
(SubgroupFeatureFlags::CLUSTERED, OpClasses::CLUSTERED),
(SubgroupFeatureFlags::QUAD, OpClasses::QUAD),
(SubgroupFeatureFlags::ROTATE, OpClasses::ROTATE),
(
SubgroupFeatureFlags::ROTATE_CLUSTERED,
OpClasses::ROTATE_CLUSTERED,
),
(SubgroupFeatureFlags::PARTITIONED_NV, OpClasses::PARTITIONED),
] {
if supported.contains(flag) {
ops |= class;
}
}
let mut arith = Arith::NONE;
if let Some(f) = physical.shader_arithmetic_features() {
if f.shader_float16 {
arith |= Arith::FLOAT16;
}
if f.shader_int8 {
arith |= Arith::INT8;
}
if f.storage_buffer_16bit {
arith |= Arith::STORAGE16;
}
if f.storage_buffer_8bit {
arith |= Arith::STORAGE8;
}
}
if physical
.shader_integer_dot_product_properties()
.is_some_and(|d| d.has_any_int8_acceleration())
{
arith |= Arith::DOT8;
}
let mut coop: Vec<CoopShape> = physical
.cooperative_matrix_properties()
.iter()
.map(|p| CoopShape {
m: p.m_size(),
n: p.n_size(),
k: p.k_size(),
a: component(p.a_type() as u32),
b: component(p.b_type() as u32),
c: component(p.c_type() as u32),
result: component(p.result_type() as u32),
saturating: p.saturating_accumulation(),
})
.collect();
coop.sort();
coop.dedup();
Some(Self {
default_subgroup: sg.subgroup_size,
subgroup_range: sg
.size_control
.map(|s| (s.min_subgroup_size, s.max_subgroup_size)),
ops,
arith,
coop,
})
}
pub fn admissible_subgroups(&self) -> Vec<Subgroup> {
let mut out = vec![Subgroup::Dynamic];
match self.subgroup_range {
Some((min, max)) => {
let mut w = min.max(1).next_power_of_two();
while w <= max {
out.push(Subgroup::Fixed(w));
match w.checked_mul(2) {
Some(next) => w = next,
None => break,
}
}
}
None => out.push(Subgroup::Fixed(self.default_subgroup)),
}
out
}
pub fn target_for(&self, subgroup: Subgroup) -> VulkanTarget {
VulkanTarget {
subgroup,
ops: self.ops,
arith: self.arith,
coop: CoopMatrix::from_shapes(self.coop.clone()),
}
}
pub fn admits(&self, target: &VulkanTarget) -> bool {
let width_ok = match target.subgroup {
Subgroup::Dynamic => true,
Subgroup::Fixed(w) => self.admissible_subgroups().contains(&Subgroup::Fixed(w)),
};
let coop_ok = match &target.coop {
CoopMatrix::None => true,
CoopMatrix::Shapes(s) => s.iter().all(|x| self.coop.contains(x)),
CoopMatrix::Digest(_) => {
matches!(
CoopMatrix::from_shapes(self.coop.clone()),
CoopMatrix::Digest(d) if CoopMatrix::Digest(d) == target.coop
)
}
};
width_ok && self.ops.contains(target.ops) && self.arith.contains(target.arith) && coop_ok
}
}
fn component(raw: u32) -> ComponentType {
match raw {
0 => ComponentType::F16,
1 => ComponentType::F32,
2 => ComponentType::F64,
3 => ComponentType::S8,
4 => ComponentType::S16,
5 => ComponentType::S32,
6 => ComponentType::S64,
7 => ComponentType::U8,
8 => ComponentType::U16,
9 => ComponentType::U32,
10 => ComponentType::U64,
n => ComponentType::Other(n),
}
}