tract_gpu/
turn_handler.rs1use crate::memory::DeviceMemSchema;
2use crate::memory::DeviceMemoryPool;
3use crate::tensor::DeviceTensor;
4use tract_core::internal::*;
5
6#[derive(Debug, Clone)]
7pub struct DeviceTurnHandler {
8 pub mem_schema: DeviceMemSchema,
9}
10
11impl DeviceTurnHandler {
12 pub fn from_plan(plan: &TypedSimplePlan, memory_hint: &SymbolValues) -> TractResult<Self> {
13 let mem_schema =
14 DeviceMemSchema::build(plan.model(), plan.order_without_consts(), memory_hint)?;
15 Ok(Self { mem_schema })
16 }
17}
18
19impl TurnStateHandler for DeviceTurnHandler {
20 fn before_plan_eval(&self, turn: &mut TurnState) -> TractResult<()> {
21 let resolved_mem_schema = self.mem_schema.resolve(&turn.resolved_symbols)?;
22 let memory_pool = DeviceMemoryPool::from_schema(resolved_mem_schema)?;
23
24 turn.shared.insert(memory_pool);
25 ensure!(turn.shared.get::<DeviceMemoryPool>().is_some());
26 Ok(())
27 }
28
29 fn after_plan_eval(&self, turn: &mut TurnState) -> TractResult<()> {
30 turn.shared.remove::<DeviceMemoryPool>();
31 Ok(())
32 }
33}
34
35pub fn make_tensor_for_node(
36 ctx: &EvalContext,
37 dt: DatumType,
38 shape: &[usize],
39) -> TractResult<DeviceTensor> {
40 ctx.shared
41 .and_then(|s| s.get::<DeviceMemoryPool>())
42 .map(|mem| mem.tensor_for_node(ctx.node_id, dt, shape))
43 .unwrap_or_else(|| DeviceTensor::uninitialized_dt(dt, shape))
44}
45
46pub fn make_scalar_exotic_tensor_for_node(
47 ctx: &EvalContext,
48 dt: DatumType,
49 exotic_fact: Box<dyn ExoticFact>,
50) -> TractResult<DeviceTensor> {
51 match ctx.shared.and_then(|s| s.get::<DeviceMemoryPool>()) {
52 Some(mem) => mem.scalar_exotic_tensor_for_node(ctx.node_id, dt, exotic_fact),
53 None => DeviceTensor::uninitialized_exotic(exotic_fact),
54 }
55}