use serde::{Deserialize, Serialize};
use std::collections::{HashMap, HashSet};
use thiserror::Error;
use vyre_foundation::{execution_plan::fusion::merge_programs_shared, ir::Program};
use vyre_lower::{KernelDescriptor, MemoryClass};
use crate::{
Artifact, ArtifactAbi, ArtifactEnvelope, ArtifactNodeId, CompileError, FusionGroupId,
FusionRecord, ResourceLifetime, TargetEntryPoint, TargetPayload, TargetPayloadFormat,
TargetProfile, TargetResourceAccess, TargetResourceBinding, TargetResourceMemory,
};
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct SelectedModule {
pub group: FusionGroupId,
pub stage: u32,
pub nodes: Vec<ArtifactNodeId>,
pub programs: Vec<Program>,
}
#[derive(Clone, Debug)]
pub struct SelectedLowering {
pub artifact: crate::Digest,
pub group: FusionGroupId,
pub stage: u32,
pub nodes: Vec<ArtifactNodeId>,
pub descriptor: KernelDescriptor,
pub abi: ArtifactAbi,
pub canonical_bindings: Vec<TargetResourceBinding>,
pub logical_element_count: u32,
program: Program,
}
pub const TARGET_MODULE_BUNDLE_SCHEMA_VERSION: u16 = 2;
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct TargetModuleImage {
pub group: FusionGroupId,
pub stage: u32,
pub nodes: Vec<ArtifactNodeId>,
pub program: Vec<u8>,
pub descriptor: KernelDescriptor,
pub entry_point: String,
pub bytes: Vec<u8>,
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct TargetModuleBundle {
pub schema_version: u16,
pub modules: Vec<TargetModuleImage>,
}
impl TargetModuleBundle {
#[must_use]
pub fn new(mut modules: Vec<TargetModuleImage>) -> Self {
modules.sort_by_key(|module| (module.stage, module.group));
Self {
schema_version: TARGET_MODULE_BUNDLE_SCHEMA_VERSION,
modules,
}
}
pub fn to_bytes(&self) -> Result<Vec<u8>, TargetCompileError> {
let body = serde_json::to_vec(self)
.map_err(|error| TargetCompileError::ModuleBundle(error.to_string()))?;
let digest = blake3::hash(&body);
let mut bytes = Vec::with_capacity(32 + body.len());
bytes.extend_from_slice(digest.as_bytes());
bytes.extend_from_slice(&body);
Ok(bytes)
}
pub fn from_bytes(bytes: &[u8]) -> Result<Self, TargetCompileError> {
let (expected, body) = bytes.split_at_checked(32).ok_or_else(|| {
TargetCompileError::ModuleBundle("target module bundle is truncated".to_string())
})?;
let actual = blake3::hash(body);
if actual.as_bytes() != expected {
return Err(TargetCompileError::ModuleBundle(
"target module bundle digest mismatch".to_string(),
));
}
let bundle: Self = serde_json::from_slice(body)
.map_err(|error| TargetCompileError::ModuleBundle(error.to_string()))?;
for module in &bundle.modules {
if module.nodes.is_empty() {
return Err(TargetCompileError::ModuleBundle(format!(
"fusion group {} has no selected nodes",
module.group.0
)));
}
Program::from_wire(&module.program).map_err(|error| {
TargetCompileError::ModuleBundle(format!(
"fusion group {} selected Program is malformed: {error}",
module.group.0
))
})?;
vyre_lower::verify_descriptor(&module.descriptor).map_err(|error| {
TargetCompileError::ModuleBundle(format!(
"fusion group {} descriptor is invalid: {error:?}",
module.group.0
))
})?;
}
if bundle.schema_version != TARGET_MODULE_BUNDLE_SCHEMA_VERSION {
return Err(TargetCompileError::ModuleBundle(format!(
"schema {} is unsupported; expected {}",
bundle.schema_version, TARGET_MODULE_BUNDLE_SCHEMA_VERSION
)));
}
if bundle.modules.windows(2).any(|modules| {
(modules[0].stage, modules[0].group) >= (modules[1].stage, modules[1].group)
}) {
return Err(TargetCompileError::ModuleBundle(
"module bundle is not in canonical stage/group order".to_string(),
));
}
let canonical = bundle.to_bytes()?;
if canonical != bytes {
return Err(TargetCompileError::ModuleBundle(
"module bundle is not in canonical stage/group order".to_string(),
));
}
Ok(bundle)
}
}
#[derive(Debug, Error)]
#[non_exhaustive]
pub enum TargetCompileError {
#[error("target compiler rejected the neutral artifact: {0}")]
InvalidArtifact(String),
#[error("target capability rejected the selected plan: {0}")]
Unsupported(String),
#[error("target emission failed: {0}")]
Emission(String),
#[error("target module bundle failed: {0}")]
ModuleBundle(String),
#[error("target payload construction failed: {0}")]
Payload(#[from] CompileError),
}
pub trait TargetCompiler: Send + Sync {
fn format(&self) -> &TargetPayloadFormat;
fn profile(&self) -> &TargetProfile;
fn compile(&self, artifact: &Artifact) -> Result<TargetPayload, TargetCompileError>;
}
pub fn attach_target(
artifact: Artifact,
compiler: &dyn TargetCompiler,
) -> Result<ArtifactEnvelope, TargetCompileError> {
let payload = compiler.compile(&artifact)?;
let mut envelope = ArtifactEnvelope::new(artifact);
envelope.attach_target_payload(payload)?;
Ok(envelope)
}
pub(crate) fn selected_modules(
artifact: &Artifact,
) -> Result<Vec<SelectedModule>, TargetCompileError> {
artifact
.fusion()
.iter()
.map(|group| decode_group(artifact, group))
.collect()
}
fn fuse_selected_module(module: &SelectedModule) -> Result<Program, TargetCompileError> {
merge_programs_shared(&module.programs).map_err(|error| {
TargetCompileError::Unsupported(format!(
"fusion group {} cannot form one target module: {error}",
module.group.0
))
})
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct EmittedTargetModule {
pub entry_point: String,
pub grid_size: [u32; 3],
pub dynamic_shared_bytes: u32,
pub workgroup_size: [u32; 3],
pub resource_bindings: Vec<TargetResourceBinding>,
pub bytes: Vec<u8>,
}
pub fn compile_selected_modules(
artifact: &Artifact,
format: TargetPayloadFormat,
profile: TargetProfile,
mut emit: impl FnMut(
&SelectedLowering,
&TargetProfile,
) -> Result<EmittedTargetModule, TargetCompileError>,
) -> Result<TargetPayload, TargetCompileError> {
let modules = selected_modules(artifact)?;
let mut images = Vec::with_capacity(modules.len());
let mut entries = Vec::with_capacity(modules.len());
for module in modules {
let program = fuse_selected_module(&module)?;
let lowered = vyre_lower::lower_verified(&program).map_err(|error| {
TargetCompileError::Emission(format!(
"verified lowering failed for fusion group {}: {error}",
module.group.0
))
})?;
let bindings = selected_resource_bindings(artifact, &module, &lowered.descriptor)?;
let abi = selected_abi(artifact, &module);
let logical_element_count =
selected_logical_element_count(artifact, &module, &lowered.program);
let selected = SelectedLowering {
artifact: artifact.digest(),
group: module.group,
stage: module.stage,
nodes: module.nodes,
descriptor: lowered.descriptor,
abi,
canonical_bindings: bindings,
logical_element_count,
program: lowered.program,
};
let emitted = emit(&selected, &profile)?;
let node = *selected.nodes.first().ok_or_else(|| {
TargetCompileError::InvalidArtifact(format!(
"fusion group {} has no member node",
selected.group.0
))
})?;
let entry_point = emitted.entry_point;
entries.push(TargetEntryPoint {
name: entry_point.clone(),
node,
workgroup_size: emitted.workgroup_size,
grid_size: emitted.grid_size,
dynamic_shared_bytes: emitted.dynamic_shared_bytes,
resource_bindings: emitted.resource_bindings,
});
let program = selected.program.to_wire().map_err(|error| {
TargetCompileError::ModuleBundle(format!(
"fusion group {} selected Program encoding failed: {error}",
selected.group.0
))
})?;
images.push(TargetModuleImage {
group: selected.group,
stage: selected.stage,
nodes: selected.nodes.clone(),
program,
descriptor: selected.descriptor.clone(),
entry_point,
bytes: emitted.bytes,
});
}
let bytes = TargetModuleBundle::new(images).to_bytes()?;
TargetPayload::new(artifact, format, profile, entries, bytes).map_err(Into::into)
}
fn selected_resource_bindings(
artifact: &Artifact,
module: &SelectedModule,
descriptor: &KernelDescriptor,
) -> Result<Vec<TargetResourceBinding>, TargetCompileError> {
let canonical_by_name = module
.nodes
.iter()
.filter_map(|node| {
artifact
.abi()
.entries
.iter()
.find(|entry| entry.node == *node)
})
.flat_map(|entry| entry.inputs.iter().chain(entry.outputs.iter()).copied())
.filter_map(|value| {
artifact
.resources()
.iter()
.find(|resource| resource.value == value)
.map(|resource| (resource.name.as_str(), value))
})
.collect::<HashMap<_, _>>();
let constant_values = artifact
.resources()
.iter()
.filter(|resource| resource.lifetime == ResourceLifetime::Constant)
.map(|resource| resource.value)
.collect::<HashSet<_>>();
descriptor
.bindings
.slots
.iter()
.filter(|slot| {
!matches!(
slot.memory_class,
MemoryClass::Shared | MemoryClass::Scratch
) && slot.name != vyre_lower::TRAP_SIDECAR_NAME
})
.map(|slot| {
let resource = canonical_by_name
.get(slot.name.as_str())
.copied()
.ok_or_else(|| {
TargetCompileError::InvalidArtifact(format!(
"fusion group {} descriptor binding `{}` has no canonical artifact resource",
module.group.0, slot.name
))
})?;
Ok(TargetResourceBinding {
resource,
group: if matches!(slot.memory_class, MemoryClass::Uniform) {
1
} else {
0
},
slot: slot.slot,
memory: if matches!(
slot.memory_class,
MemoryClass::Constant | MemoryClass::Uniform
) || constant_values.contains(&resource)
{
TargetResourceMemory::Constant
} else {
TargetResourceMemory::Global
},
access: match slot.visibility {
vyre_lower::BindingVisibility::ReadOnly => TargetResourceAccess::ReadOnly,
vyre_lower::BindingVisibility::WriteOnly => TargetResourceAccess::WriteOnly,
vyre_lower::BindingVisibility::ReadWrite => TargetResourceAccess::ReadWrite,
},
})
})
.collect()
}
fn selected_logical_element_count(
artifact: &Artifact,
module: &SelectedModule,
program: &Program,
) -> u32 {
let nodes = module.nodes.iter().copied().collect::<HashSet<_>>();
let values = artifact
.abi()
.entries
.iter()
.filter(|entry| nodes.contains(&entry.node))
.flat_map(|entry| entry.inputs.iter().chain(&entry.outputs))
.copied()
.collect::<HashSet<_>>();
let full_span = program.stats().atomic_op_count > 0
|| vyre_foundation::program_caps::scan(program).subgroup_ops;
let selected = artifact
.resources()
.iter()
.filter(|resource| values.contains(&resource.value));
let count = if full_span {
selected.map(|resource| resource.element_count).max()
} else {
selected
.filter(|resource| {
artifact
.abi()
.resources
.iter()
.find(|abi| abi.value == resource.value)
.is_some_and(|abi| {
matches!(
abi.access,
crate::AbiAccess::WriteOnly | crate::AbiAccess::ReadWrite
)
})
})
.map(|resource| resource.element_count)
.max()
.or_else(|| {
artifact
.resources()
.iter()
.filter(|resource| values.contains(&resource.value))
.map(|resource| resource.element_count)
.max()
})
}
.unwrap_or(1)
.max(1);
u32::try_from(count).unwrap_or(u32::MAX)
}
fn selected_abi(artifact: &Artifact, module: &SelectedModule) -> ArtifactAbi {
let nodes = module.nodes.iter().copied().collect::<HashSet<_>>();
let entries = artifact
.abi()
.entries
.iter()
.filter(|entry| nodes.contains(&entry.node))
.cloned()
.collect::<Vec<_>>();
let values = entries
.iter()
.flat_map(|entry| entry.inputs.iter().chain(&entry.outputs))
.copied()
.collect::<HashSet<_>>();
ArtifactAbi {
resources: artifact
.abi()
.resources
.iter()
.filter(|resource| values.contains(&resource.value))
.cloned()
.collect(),
entries,
}
}
fn decode_group(
artifact: &Artifact,
group: &FusionRecord,
) -> Result<SelectedModule, TargetCompileError> {
let mut nodes = group.members.clone();
nodes.sort();
let programs = nodes
.iter()
.map(|node| {
let record = artifact
.nodes()
.iter()
.find(|record| record.id == *node)
.ok_or_else(|| {
TargetCompileError::InvalidArtifact(format!(
"fusion group {} references missing node {}",
group.id.0, node.0
))
})?;
Program::from_wire(&record.program).map_err(|error| {
TargetCompileError::InvalidArtifact(format!(
"node {} canonical Program failed to decode: {error}",
node.0
))
})
})
.collect::<Result<Vec<_>, _>>()?;
Ok(SelectedModule {
group: group.id,
stage: group.stage,
nodes,
programs,
})
}