use alloc::vec::Vec;
use core::ops::Deref;
use cubecl_environment::sync::{Arc, Mutex};
use crate::device_events::EventApi;
use crate::driver::DriverError;
pub struct Event<A: EventApi> {
sys: A::Event,
}
impl<A: EventApi> Event<A> {
pub fn new() -> Result<Self, DriverError> {
Ok(Self {
sys: A::event_create()?,
})
}
pub fn record(&self, stream: A::Stream) -> Result<(), DriverError> {
A::event_record(&self.sys, stream)
}
pub fn wait(&self) -> Result<(), DriverError> {
A::event_wait(&self.sys)
}
pub fn wait_async(&self, stream: A::Stream) -> Result<(), DriverError> {
A::stream_wait_event(stream, &self.sys)
}
pub fn elapsed(&self, other: &Self) -> Result<cubecl_common::profile::Duration, DriverError> {
A::event_elapsed(&self.sys, &other.sys)
}
}
impl<A: EventApi> Drop for Event<A> {
fn drop(&mut self) {
if let Err(err) = A::event_destroy(&mut self.sys) {
log::warn!("Failed to release a {} event: {err}", A::BACKEND);
}
}
}
impl<A: EventApi> core::fmt::Debug for Event<A> {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
write!(f, "{}Event", A::BACKEND)
}
}
pub(crate) struct EventPool<A: EventApi> {
free: Arc<Mutex<Vec<Event<A>>>>,
}
impl<A: EventApi> EventPool<A> {
pub(crate) fn acquire(&self) -> Result<Pooled<A>, DriverError> {
let event = match self.free.lock().pop() {
Some(event) => event,
None => Event::new()?,
};
Ok(Pooled {
event: Some(event),
pool: self.clone(),
})
}
fn release(&self, event: Event<A>) {
self.free.lock().push(event);
}
}
impl<A: EventApi> Clone for EventPool<A> {
fn clone(&self) -> Self {
Self {
free: self.free.clone(),
}
}
}
impl<A: EventApi> Default for EventPool<A> {
fn default() -> Self {
Self {
free: Arc::new(Mutex::new(Vec::new())),
}
}
}
impl<A: EventApi> core::fmt::Debug for EventPool<A> {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
write!(f, "EventPool<{}>({})", A::BACKEND, self.free.lock().len())
}
}
pub(crate) struct Pooled<A: EventApi> {
event: Option<Event<A>>,
pool: EventPool<A>,
}
impl<A: EventApi> Deref for Pooled<A> {
type Target = Event<A>;
fn deref(&self) -> &Self::Target {
self.event.as_ref().expect("emptied only by Drop")
}
}
impl<A: EventApi> Drop for Pooled<A> {
fn drop(&mut self) {
if let Some(event) = self.event.take() {
self.pool.release(event);
}
}
}
impl<A: EventApi> core::fmt::Debug for Pooled<A> {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
write!(f, "Pooled<{}>", A::BACKEND)
}
}