Skip to main content

tvm_ffi/
device.rs

1/*
2 * Licensed to the Apache Software Foundation (ASF) under one
3 * or more contributor license agreements.  See the NOTICE file
4 * distributed with this work for additional information
5 * regarding copyright ownership.  The ASF licenses this file
6 * to you under the Apache License, Version 2.0 (the
7 * "License"); you may not use this file except in compliance
8 * with the License.  You may obtain a copy of the License at
9 *
10 *   http://www.apache.org/licenses/LICENSE-2.0
11 *
12 * Unless required by applicable law or agreed to in writing,
13 * software distributed under the License is distributed on an
14 * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
15 * KIND, either express or implied.  See the License for the
16 * specific language governing permissions and limitations
17 * under the License.
18 */
19use crate::error::Result;
20use crate::type_traits::AnyCompatible;
21use tvm_ffi_sys::dlpack::DLDevice;
22use tvm_ffi_sys::{TVMFFIAny, TVMFFITypeIndex as TypeIndex};
23use tvm_ffi_sys::{TVMFFIEnvGetStream, TVMFFIEnvSetStream, TVMFFIStreamHandle};
24
25/// Get the current stream for a device
26pub fn current_stream(device: &DLDevice) -> TVMFFIStreamHandle {
27    unsafe { TVMFFIEnvGetStream(device.device_type as i32, device.device_id) }
28}
29/// Call `f` with the device stream temporarily set to `stream`.
30///
31/// The previous stream is restored however `f` exits: when it returns a
32/// value, when it returns an error, and when it panics.
33///
34/// # Safety
35///
36/// `stream` must be a valid stream handle for the given device, or null.
37pub unsafe fn with_stream<T>(
38    device: &DLDevice,
39    stream: TVMFFIStreamHandle,
40    f: impl FnOnce() -> Result<T>,
41) -> Result<T> {
42    let mut prev_stream: TVMFFIStreamHandle = std::ptr::null_mut();
43    unsafe {
44        crate::check_safe_call!(TVMFFIEnvSetStream(
45            device.device_type as i32,
46            device.device_id,
47            stream,
48            &mut prev_stream as *mut TVMFFIStreamHandle
49        ))?;
50    }
51    let restore = RestoreStream {
52        device: *device,
53        stream: prev_stream,
54    };
55    let result = f();
56    let restored = restore.restore();
57    let value = result?;
58    restored?;
59    Ok(value)
60}
61
62/// Restores a device's previous stream when dropped, so that an unwinding
63/// panic in `with_stream`'s closure still restores it.
64struct RestoreStream {
65    device: DLDevice,
66    stream: TVMFFIStreamHandle,
67}
68
69impl RestoreStream {
70    fn set(&self) -> i32 {
71        unsafe {
72            TVMFFIEnvSetStream(
73                self.device.device_type as i32,
74                self.device.device_id,
75                self.stream,
76                std::ptr::null_mut(),
77            )
78        }
79    }
80
81    /// Restores the stream now, reporting a failure.
82    fn restore(self) -> Result<()> {
83        let restore = std::mem::ManuallyDrop::new(self);
84        crate::check_safe_call!(restore.set())
85    }
86}
87
88impl Drop for RestoreStream {
89    fn drop(&mut self) {
90        // Only reached while unwinding; a failure cannot be reported there.
91        self.set();
92    }
93}
94
95/// AnyCompatible for DLDevice
96unsafe impl AnyCompatible for DLDevice {
97    fn type_str() -> String {
98        // make it consistent with c++ representation
99        "Device".to_string()
100    }
101
102    unsafe fn copy_to_any_view(src: &Self, data: &mut TVMFFIAny) {
103        data.type_index = TypeIndex::kTVMFFIDevice as i32;
104        data.small_str_len = 0;
105        data.data_union.v_uint64 = 0;
106        data.data_union.v_device = *src;
107    }
108
109    unsafe fn move_to_any(src: Self, data: &mut TVMFFIAny) {
110        data.type_index = TypeIndex::kTVMFFIDevice as i32;
111        data.small_str_len = 0;
112        data.data_union.v_int64 = 0;
113        data.data_union.v_device = src;
114    }
115
116    unsafe fn check_any_strict(data: &TVMFFIAny) -> bool {
117        return data.type_index == TypeIndex::kTVMFFIDevice as i32;
118    }
119
120    unsafe fn copy_from_any_view_after_check(data: &TVMFFIAny) -> Self {
121        data.data_union.v_device
122    }
123
124    unsafe fn move_from_any_after_check(data: &mut TVMFFIAny) -> Self {
125        data.data_union.v_device
126    }
127
128    unsafe fn try_cast_from_any_view(data: &TVMFFIAny) -> Result<Self, ()> {
129        if data.type_index == TypeIndex::kTVMFFIDevice as i32 {
130            Ok(data.data_union.v_device)
131        } else {
132            Err(())
133        }
134    }
135}