use serde::{Deserialize, Serialize};
pub const DISCOVERY_SCHEMA_VERSION: u32 = 1;
#[derive(Debug, Clone, Eq, PartialEq, Serialize, Deserialize)]
#[serde(tag = "status", content = "reasons", rename_all = "snake_case")]
pub enum DescriptionCompleteness {
Complete,
Partial(Vec<String>),
Unsupported(Vec<String>),
}
#[derive(Debug, Clone, Eq, PartialEq, Serialize, Deserialize)]
#[serde(tag = "kind", content = "value", rename_all = "snake_case")]
pub enum SymbolicDimension {
Known(usize),
Batch,
Sequence,
TokenRows,
Context,
MediaPositions,
Unknown,
}
#[derive(Debug, Clone, Eq, PartialEq, Serialize, Deserialize)]
pub struct TensorAxis {
pub name: String,
pub dimension: SymbolicDimension,
}
#[derive(Debug, Clone, Copy, Eq, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum ArchitectureNodeKind {
Embedding,
DecoderBlock,
Normalization,
Attention,
Mixer,
FeedForward,
MixtureOfExperts,
Router,
RoutedExperts,
SharedExperts,
ResidualAdd,
Sum,
OutputHead,
Processor,
Encoder,
Projector,
ModalityMerge,
Prediction,
Realtime,
Opaque,
}
#[derive(Debug, Clone, Copy, Eq, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum HeadSharing {
MultiHead,
MultiQuery,
GroupedQuery,
}
#[derive(Debug, Clone, Eq, PartialEq, Serialize, Deserialize)]
#[serde(tag = "kind", rename_all = "snake_case")]
pub enum ReceptiveField {
Full,
Sliding {
window: usize,
},
Local {
window: Option<usize>,
},
}
#[derive(Debug, Clone, Copy, Eq, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum AttentionMechanism {
Softmax,
Linear,
Latent,
CompressedSparse,
}
#[derive(Debug, Clone, Copy, Eq, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum PositionalEncoding {
Rotary,
Relative,
Learned,
None,
}
#[derive(Debug, Clone, Default, Eq, PartialEq, Serialize, Deserialize)]
pub struct AttentionAttributes {
pub head_sharing: Option<HeadSharing>,
pub query_heads: Option<usize>,
pub key_value_heads: Option<usize>,
pub key_head_dimension: Option<usize>,
pub value_head_dimension: Option<usize>,
pub receptive_field: Option<ReceptiveField>,
pub mechanism: Option<AttentionMechanism>,
pub recurrent: Option<bool>,
pub causal: Option<bool>,
pub positional_encoding: Option<PositionalEncoding>,
}
#[derive(Debug, Clone, Copy, Eq, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum MixerMechanism {
GatedDelta,
SelectiveStateSpace,
ShortConvolution,
}
#[derive(Debug, Clone, Eq, PartialEq, Serialize, Deserialize)]
pub struct MixerAttributes {
pub mechanism: MixerMechanism,
pub recurrent: bool,
pub convolution_width: Option<usize>,
}
#[derive(Debug, Clone, Copy, Eq, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum RoutingGranularity {
Token,
Sequence,
}
#[derive(Debug, Clone, Copy, Eq, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum RoutingScoreTransform {
Softmax,
SelectedSoftmax,
Sigmoid,
SqrtSoftplus,
Identity,
}
#[derive(Debug, Clone, Copy, Eq, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum RoutingNormalization {
SelectedSum,
SelectedAndSharedSum,
None,
}
#[derive(Debug, Clone, Eq, PartialEq, Serialize, Deserialize)]
pub struct MoeAttributes {
pub routed_experts: usize,
pub selected_experts: usize,
pub shared_experts: Option<usize>,
pub shared_expert_width: Option<usize>,
pub shared_expert_gated: Option<bool>,
pub granularity: Option<RoutingGranularity>,
pub score_transform: Option<RoutingScoreTransform>,
pub normalization: Option<RoutingNormalization>,
}
#[derive(Debug, Clone, Eq, PartialEq, Serialize, Deserialize)]
pub struct ArchitectureNode {
pub id: String,
pub label: String,
pub kind: ArchitectureNodeKind,
pub parent: Option<String>,
pub layer_index: Option<usize>,
pub parameter_groups: Vec<String>,
pub observation_paths: Vec<String>,
pub output_axes: Option<Vec<TensorAxis>>,
pub attention: Option<AttentionAttributes>,
pub mixer: Option<MixerAttributes>,
pub moe: Option<MoeAttributes>,
pub completeness: DescriptionCompleteness,
}
#[derive(Debug, Clone, Copy, Eq, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum ArchitectureEdgeKind {
Data,
Residual,
Routing,
State,
}
#[derive(Debug, Clone, Eq, PartialEq, Serialize, Deserialize)]
pub struct ArchitectureEdge {
pub from: String,
pub to: String,
pub kind: ArchitectureEdgeKind,
}
#[derive(Debug, Clone, Eq, PartialEq, Serialize, Deserialize)]
pub struct ArchitectureParameterGroup {
pub id: String,
pub canonical_prefix: String,
}
#[derive(Debug, Clone, Eq, PartialEq, Serialize, Deserialize)]
pub struct ArchitectureDescriptor {
pub schema_version: u32,
pub nodes: Vec<ArchitectureNode>,
pub edges: Vec<ArchitectureEdge>,
pub parameter_groups: Vec<ArchitectureParameterGroup>,
pub observations: ObservationCatalog,
pub completeness: DescriptionCompleteness,
}
impl ArchitectureDescriptor {
pub fn node(&self, id: &str) -> Option<&ArchitectureNode> {
self.nodes.iter().find(|node| node.id == id)
}
}
#[derive(Debug, Clone, Copy, Eq, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum UnitObservation {
Input,
Output,
}
impl UnitObservation {
pub fn path(self, unit: &str) -> String {
format!(
"{unit}.{}",
match self {
Self::Input => "input",
Self::Output => "output",
}
)
}
}
#[derive(Debug, Clone, Copy, Eq, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum RoutingObservationField {
SelectedExperts,
SelectedScores,
Coefficients,
RoutedOutput,
LocalRoutedOutput,
ReducedRoutedOutput,
SharedOutput,
CombinedOutput,
}
impl RoutingObservationField {
pub fn path(self, module: &str) -> String {
let field = match self {
Self::SelectedExperts => "selected_experts",
Self::SelectedScores => "selected_scores",
Self::Coefficients => "coefficients",
Self::RoutedOutput => "routed_output",
Self::LocalRoutedOutput => "local_routed_output",
Self::ReducedRoutedOutput => "reduced_routed_output",
Self::SharedOutput => "shared_output",
Self::CombinedOutput => "combined_output",
};
format!("{module}.routing.{field}")
}
}
#[derive(Debug, Clone, Copy, Eq, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum ObservationDtype {
Floating,
Integer,
Boolean,
Unknown,
}
#[derive(Debug, Clone, Copy, Eq, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum ObservationValueType {
Tensor,
}
#[derive(Debug, Clone, Copy, Eq, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum ObservationPosition {
BeforeIntervention,
ReadOnly,
AfterIntervention,
}
#[derive(Debug, Clone, Copy, Eq, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum ObservationRequirement {
ActivationHooks,
RoutingEvents,
MediaInput,
PredictionExecution,
}
#[derive(Debug, Clone, Eq, PartialEq, Serialize, Deserialize)]
pub struct ObservationPoint {
pub path: String,
pub node_id: String,
pub meaning: String,
pub value_type: ObservationValueType,
pub dtype: ObservationDtype,
pub axes: Option<Vec<TensorAxis>>,
pub prefill: bool,
pub decode: bool,
pub requirements: Vec<ObservationRequirement>,
pub position: ObservationPosition,
pub retained_bytes: Option<u64>,
pub host_bytes: Option<u64>,
}
#[derive(Debug, Clone, Eq, PartialEq, Serialize, Deserialize)]
pub struct ObservationCatalog {
pub schema_version: u32,
pub points: Vec<ObservationPoint>,
pub completeness: DescriptionCompleteness,
}
impl ObservationCatalog {
pub fn get(&self, path: &str) -> Option<&ObservationPoint> {
self.points.iter().find(|point| point.path == path)
}
}
#[derive(Debug, Clone, Copy, Default, Eq, PartialEq, Serialize, Deserialize)]
pub struct ObservationMechanisms {
pub activation_tensors: bool,
pub routing_tensors: bool,
pub floating_to_f32: bool,
}
#[derive(Debug, Clone, Eq, PartialEq, Serialize, Deserialize)]
#[serde(tag = "status", content = "reason", rename_all = "snake_case")]
pub enum ObservationSupportStatus {
Supported,
Conditional(String),
Unsupported(String),
Unverified(String),
}
#[derive(Debug, Clone, Eq, PartialEq, Serialize, Deserialize)]
pub struct ObservationSupport {
pub path: String,
pub prefill: ObservationSupportStatus,
pub decode: ObservationSupportStatus,
pub floating_to_f32: bool,
}
#[derive(Debug, Clone, Eq, PartialEq, Serialize, Deserialize)]
pub struct ObservationSupportReport {
pub schema_version: u32,
pub points: Vec<ObservationSupport>,
#[serde(default)]
pub capture: crate::capture::CaptureCapabilities,
}