use std::collections::HashMap;
use crate::error::{GraphError, GraphResult};
use crate::graph::ComputeGraph;
use crate::node::{GraphNode, MemcpyDir, NodeId, NodeKind, StreamId};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CaptureStatus {
None,
Active,
Invalidated,
}
impl std::fmt::Display for CaptureStatus {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::None => write!(f, "none"),
Self::Active => write!(f, "active"),
Self::Invalidated => write!(f, "invalidated"),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct CaptureEvent(pub u32);
impl std::fmt::Display for CaptureEvent {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "E{}", self.0)
}
}
#[derive(Debug, Default)]
pub struct StreamCapture {
graph: ComputeGraph,
status: HashMap<StreamId, CaptureStatus>,
last_on_stream: HashMap<StreamId, NodeId>,
event_source: HashMap<CaptureEvent, NodeId>,
pending_waits: HashMap<StreamId, Vec<NodeId>>,
poisoned: bool,
}
impl StreamCapture {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn status(&self, stream: StreamId) -> CaptureStatus {
self.status
.get(&stream)
.copied()
.unwrap_or(CaptureStatus::None)
}
#[must_use]
pub fn is_capturing(&self) -> bool {
self.status.values().any(|s| *s == CaptureStatus::Active)
}
pub fn begin_capture(&mut self, stream: StreamId) -> GraphResult<()> {
match self.status(stream) {
CaptureStatus::Active => Err(GraphError::Internal(format!(
"stream {stream} is already capturing"
))),
_ => {
self.status.insert(stream, CaptureStatus::Active);
Ok(())
}
}
}
pub fn record_kernel(
&mut self,
stream: StreamId,
function_name: &str,
num_blocks: u32,
threads_per_block: u32,
shared_mem: u32,
) -> GraphResult<NodeId> {
let kind = NodeKind::KernelLaunch {
function_name: function_name.to_owned(),
config: crate::node::KernelConfig::linear(num_blocks, threads_per_block, shared_mem),
fusible: false,
};
self.record_op(stream, kind, function_name)
}
pub fn record_memcpy(
&mut self,
stream: StreamId,
name: &str,
dir: MemcpyDir,
size_bytes: usize,
) -> GraphResult<NodeId> {
self.record_op(stream, NodeKind::Memcpy { dir, size_bytes }, name)
}
pub fn record_memset(
&mut self,
stream: StreamId,
name: &str,
size_bytes: usize,
value: u8,
) -> GraphResult<NodeId> {
self.record_op(stream, NodeKind::Memset { size_bytes, value }, name)
}
pub fn record_host_callback(&mut self, stream: StreamId, label: &str) -> GraphResult<NodeId> {
self.record_op(
stream,
NodeKind::HostCallback {
label: label.to_owned(),
},
label,
)
}
fn record_op(&mut self, stream: StreamId, kind: NodeKind, name: &str) -> GraphResult<NodeId> {
self.ensure_active(stream)?;
let node = GraphNode::new(NodeId(0), kind)
.with_name(name)
.with_stream(stream);
let id = self.graph.add_node(node);
self.drain_pending_waits(stream, id)?;
if let Some(&prev) = self.last_on_stream.get(&stream) {
self.graph.add_edge(prev, id)?;
}
self.last_on_stream.insert(stream, id);
Ok(id)
}
pub fn record_event(&mut self, stream: StreamId, event: CaptureEvent) -> GraphResult<()> {
self.ensure_active(stream)?;
if let Some(&head) = self.last_on_stream.get(&stream) {
self.event_source.insert(event, head);
}
Ok(())
}
pub fn wait_event(&mut self, stream: StreamId, event: CaptureEvent) -> GraphResult<()> {
self.ensure_active(stream)?;
let Some(&source) = self.event_source.get(&event) else {
return Ok(());
};
self.pending_waits.entry(stream).or_default().push(source);
Ok(())
}
pub fn end_capture(mut self, stream: StreamId) -> GraphResult<ComputeGraph> {
match self.status(stream) {
CaptureStatus::Active => {}
CaptureStatus::Invalidated => {
return Err(GraphError::Internal(
"capture was invalidated by an illegal operation".to_owned(),
));
}
CaptureStatus::None => {
return Err(GraphError::Internal(format!(
"stream {stream} is not capturing"
)));
}
}
if self.poisoned {
return Err(GraphError::Internal(
"capture was invalidated on another stream in the session".to_owned(),
));
}
self.status.insert(stream, CaptureStatus::None);
Ok(std::mem::take(&mut self.graph))
}
#[must_use]
pub fn node_count(&self) -> usize {
self.graph.node_count()
}
pub fn invalidate(&mut self, stream: StreamId) -> GraphResult<()> {
self.ensure_active(stream)?;
self.status.insert(stream, CaptureStatus::Invalidated);
self.poisoned = true;
Ok(())
}
fn ensure_active(&self, stream: StreamId) -> GraphResult<()> {
match self.status(stream) {
CaptureStatus::Active => Ok(()),
CaptureStatus::Invalidated => Err(GraphError::Internal(format!(
"stream {stream} capture is invalidated"
))),
CaptureStatus::None => Err(GraphError::Internal(format!(
"stream {stream} is not capturing"
))),
}
}
fn drain_pending_waits(&mut self, stream: StreamId, id: NodeId) -> GraphResult<()> {
if let Some(sources) = self.pending_waits.remove(&stream) {
for source in sources {
if source != id {
if let Err(e) = self.graph.add_edge(source, id) {
self.status.insert(stream, CaptureStatus::Invalidated);
self.poisoned = true;
return Err(e);
}
}
}
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::node::{MemcpyDir, StreamId};
fn s(n: u32) -> StreamId {
StreamId(n)
}
#[test]
fn fresh_stream_is_not_capturing() {
let cap = StreamCapture::new();
assert_eq!(cap.status(s(0)), CaptureStatus::None);
assert!(!cap.is_capturing());
}
#[test]
fn begin_sets_active() {
let mut cap = StreamCapture::new();
cap.begin_capture(s(0)).expect("begin capture");
assert_eq!(cap.status(s(0)), CaptureStatus::Active);
assert!(cap.is_capturing());
}
#[test]
fn double_begin_rejected() {
let mut cap = StreamCapture::new();
cap.begin_capture(s(0)).expect("begin capture");
assert!(cap.begin_capture(s(0)).is_err());
}
#[test]
fn record_on_uncaptured_stream_rejected() {
let mut cap = StreamCapture::new();
assert!(cap.record_kernel(s(0), "k", 1, 32, 0).is_err());
}
#[test]
fn sequential_ops_are_chained() {
let mut cap = StreamCapture::new();
cap.begin_capture(s(0)).expect("begin");
let a = cap.record_kernel(s(0), "a", 1, 32, 0).expect("a");
let b = cap.record_kernel(s(0), "b", 1, 32, 0).expect("b");
let c = cap.record_kernel(s(0), "c", 1, 32, 0).expect("c");
let g = cap.end_capture(s(0)).expect("end");
assert!(g.is_reachable(a, b));
assert!(g.is_reachable(b, c));
assert!(g.is_reachable(a, c));
assert!(!g.is_reachable(c, a));
assert_eq!(g.edge_count(), 2);
}
#[test]
fn fork_join_reconstructs_dependencies() {
let mut cap = StreamCapture::new();
let main = s(0);
let side = s(1);
cap.begin_capture(main).expect("begin main");
let up = cap
.record_memcpy(main, "up", MemcpyDir::HostToDevice, 1024)
.expect("up");
cap.begin_capture(side).expect("begin side");
let ev = CaptureEvent(0);
cap.record_event(main, ev).expect("record ev");
cap.wait_event(side, ev).expect("wait ev");
let k = cap.record_kernel(side, "k", 1, 32, 0).expect("k");
let ev2 = CaptureEvent(1);
cap.record_event(side, ev2).expect("record ev2");
cap.wait_event(main, ev2).expect("wait ev2");
let dn = cap
.record_memcpy(main, "dn", MemcpyDir::DeviceToHost, 1024)
.expect("dn");
let g = cap.end_capture(main).expect("end");
assert!(g.is_reachable(up, k), "fork dependency up→k missing");
assert!(g.is_reachable(k, dn), "join dependency k→dn missing");
}
#[test]
fn wait_on_unrecorded_event_is_noop() {
let mut cap = StreamCapture::new();
cap.begin_capture(s(0)).expect("begin");
assert!(cap.wait_event(s(0), CaptureEvent(9)).is_ok());
let g = cap.end_capture(s(0)).expect("end");
assert_eq!(g.edge_count(), 0);
}
#[test]
fn end_without_begin_errs() {
let cap = StreamCapture::new();
assert!(cap.end_capture(s(0)).is_err());
}
#[test]
fn pending_wait_applies_to_next_op() {
let mut cap = StreamCapture::new();
let main = s(0);
let side = s(1);
cap.begin_capture(main).expect("begin main");
cap.begin_capture(side).expect("begin side");
let a = cap.record_kernel(main, "a", 1, 32, 0).expect("a");
let ev = CaptureEvent(0);
cap.record_event(main, ev).expect("rec");
cap.wait_event(side, ev).expect("wait"); let b = cap.record_kernel(side, "b", 1, 32, 0).expect("b");
let g = cap.end_capture(main).expect("end");
assert!(g.is_reachable(a, b), "pending wait a→b not applied");
}
#[test]
fn status_display() {
assert_eq!(CaptureStatus::None.to_string(), "none");
assert_eq!(CaptureStatus::Active.to_string(), "active");
assert_eq!(CaptureStatus::Invalidated.to_string(), "invalidated");
}
#[test]
fn event_display() {
assert_eq!(CaptureEvent(3).to_string(), "E3");
}
#[test]
fn explicit_invalidation_aborts_capture() {
let mut cap = StreamCapture::new();
let main = s(0);
cap.begin_capture(main).expect("begin main");
cap.record_kernel(main, "a", 1, 32, 0).expect("a");
cap.invalidate(main).expect("invalidate");
assert_eq!(cap.status(main), CaptureStatus::Invalidated);
assert!(cap.record_kernel(main, "b", 1, 32, 0).is_err());
assert!(
cap.end_capture(main).is_err(),
"invalidated capture cannot end"
);
}
#[test]
fn invalidation_poisons_other_streams() {
let mut cap = StreamCapture::new();
let main = s(0);
let side = s(1);
cap.begin_capture(main).expect("begin main");
cap.begin_capture(side).expect("begin side");
cap.record_kernel(side, "x", 1, 32, 0).expect("x");
cap.invalidate(main).expect("invalidate main");
assert_eq!(cap.status(side), CaptureStatus::Active);
assert!(
cap.end_capture(side).is_err(),
"poisoned session cannot end"
);
}
#[test]
fn node_count_tracks_recorded_ops() {
let mut cap = StreamCapture::new();
cap.begin_capture(s(0)).expect("begin");
cap.record_kernel(s(0), "a", 1, 32, 0).expect("a");
cap.record_memset(s(0), "z", 64, 0).expect("z");
assert_eq!(cap.node_count(), 2);
}
}