use super::{DeviceStream, Driver};
use crate::id::GraphId;
use crate::memory_management::ManagedMemoryHandle;
use crate::server::{BufferBinding, ServerError};
use alloc::format;
use alloc::vec::Vec;
use core::marker::PhantomData;
use cubecl_common::bytes::Bytes;
use cubecl_environment::backtrace::BackTrace;
use cubecl_environment::collections::HashMap;
use cubecl_environment::stream::StreamId;
pub trait GraphDriver: Driver {
type Executable;
fn begin(stream: &mut Self::Stream) -> Result<(), ServerError>;
fn instantiate(
stream: &mut Self::Stream,
doomed: Option<ServerError>,
) -> Result<Self::Executable, ServerError>;
fn upload(exec: &Self::Executable, stream: &mut Self::Stream);
fn replay(exec: &Self::Executable, stream: &mut Self::Stream) -> Result<(), ServerError>;
}
pub struct Graph<D: GraphDriver> {
exec: D::Executable,
_retained: Vec<ManagedMemoryHandle>,
_retained_host: Vec<Bytes>,
written: Vec<BufferBinding>,
}
pub struct Refused {
pub error: ServerError,
pub written: Vec<BufferBinding>,
}
pub struct Captures<D: GraphDriver> {
graphs: HashMap<GraphId, Graph<D>>,
}
impl<D: GraphDriver> core::fmt::Debug for Captures<D> {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("Captures")
.field("graphs", &self.graphs.len())
.finish()
}
}
impl<D: GraphDriver> Default for Captures<D> {
fn default() -> Self {
Self {
graphs: HashMap::default(),
}
}
}
impl<D: GraphDriver> Captures<D> {
pub fn insert(&mut self, id: GraphId, graph: Graph<D>) {
self.graphs.insert(id, graph);
}
pub fn contains(&self, id: GraphId) -> bool {
self.graphs.contains_key(&id)
}
pub fn extend_written(&self, id: GraphId, written: &mut Vec<BufferBinding>) {
if let Some(graph) = self.graphs.get(&id) {
written.extend(graph.written.iter().cloned());
}
}
pub fn replay(&self, id: GraphId, stream: &mut D::Stream) -> Result<(), ServerError> {
let graph = self.graphs.get(&id).ok_or_else(|| ServerError::Generic {
reason: "replay was given an unknown or already-destroyed graph".into(),
backtrace: BackTrace::capture(),
})?;
D::replay(&graph.exec, stream)
}
pub fn destroy(&mut self, id: GraphId, stream: &mut D::Stream) {
self.graphs.remove(&id);
stream.info_cache().graph_release(id);
}
}
pub struct Window<'a, D: GraphDriver> {
stream: &'a mut D::Stream,
driver: PhantomData<D>,
}
impl<'a, D: GraphDriver> Window<'a, D> {
pub fn on(stream: &'a mut D::Stream) -> Self {
Self {
stream,
driver: PhantomData,
}
}
pub fn prepare(&mut self, stream_id: StreamId) -> Result<(), ServerError> {
self.stream.capturing().prepare(stream_id)?;
self.stream.device_memory().capture_begin();
self.stream.host_memory().capture_begin();
Ok(())
}
pub fn begin(&mut self) -> Result<(), ServerError> {
self.stream.capturing().begin()?;
let signal = self.stream.signal();
self.stream.drop_queue().drain(|| D::Stream::fence(signal));
self.stream.device_memory().capture_priming_end();
self.stream.host_memory().capture_priming_end();
if let Err(err) = D::begin(self.stream) {
self.stream.device_memory().capture_end();
self.stream.host_memory().capture_end();
self.stream.info_cache().capture_discard();
self.stream.capturing().abort();
return Err(err);
}
Ok(())
}
pub fn instantiate(&mut self, stream_id: StreamId, id: GraphId) -> Result<Graph<D>, Refused> {
let outcome = match self.stream.capturing().end(stream_id) {
Ok(outcome) => outcome,
Err(error) => {
let written = self.stream.capturing().take_recorded();
if self.stream.capturing().is_active() {
drop(self.stream.device_memory().capture_end());
drop(self.stream.host_memory().capture_end());
self.stream.info_cache().capture_discard();
self.stream.capturing().abort();
}
return Err(Refused { error, written });
}
};
let doomed = self.stream.capturing().take_failure().map(|reason| {
ServerError::graph_state(format!(
"an operation inside the capture window failed or was skipped, so the \
recording is missing an operation and cannot seal: {reason}"
))
});
let exec = D::instantiate(self.stream, doomed.clone());
let mut retained = self.stream.device_memory().capture_end();
retained.extend(self.stream.host_memory().capture_end());
let retained_host = self.stream.capturing().take_retained_host();
let signal = self.stream.signal();
self.stream.drop_queue().drain(|| D::Stream::fence(signal));
let written = self.stream.capturing().take_recorded();
let exec = match outcome.is_abandoned() {
false => exec,
true => Err(outcome.abandoned_error(stream_id, doomed)),
};
match exec {
Ok(exec) => {
self.stream.info_cache().capture_commit(id);
D::upload(&exec, self.stream);
Ok(Graph {
exec,
_retained: retained,
_retained_host: retained_host,
written,
})
}
Err(error) => {
self.stream.info_cache().capture_discard();
Err(Refused { error, written })
}
}
}
}