cubecl_core/frontend/container/sequence/
launch.rs1use alloc::vec::Vec;
2
3use cubecl_runtime::runtime::Runtime;
4use cubecl_zspace::SmallVec;
5
6use crate::{
7 compute::{KernelBuilder, KernelLauncher},
8 prelude::{CubeType, LaunchArg},
9};
10
11use super::{Sequence, SequenceExpand};
12
13pub struct SequenceArg<R: Runtime, T: LaunchArg> {
14 pub values: SmallVec<[T::RuntimeArg<R>; 5]>,
15}
16
17impl<R: Runtime, T: LaunchArg> Default for SequenceArg<R, T> {
18 fn default() -> Self {
19 Self::new()
20 }
21}
22
23impl<R: Runtime, T: LaunchArg> SequenceArg<R, T> {
24 pub fn new() -> Self {
25 Self {
26 values: SmallVec::new(),
27 }
28 }
29 pub fn push(&mut self, arg: T::RuntimeArg<R>) {
30 self.values.push(arg);
31 }
32}
33
34pub struct SequenceCompilationArg<C: LaunchArg> {
35 pub values: SmallVec<[C::CompilationArg; 5]>,
36}
37
38impl<C: LaunchArg> Clone for SequenceCompilationArg<C> {
39 fn clone(&self) -> Self {
40 Self {
41 values: self.values.clone(),
42 }
43 }
44}
45
46impl<C: LaunchArg> core::hash::Hash for SequenceCompilationArg<C> {
47 fn hash<H: core::hash::Hasher>(&self, state: &mut H) {
48 self.values.hash(state)
49 }
50}
51
52impl<C: LaunchArg> core::cmp::PartialEq for SequenceCompilationArg<C> {
53 fn eq(&self, other: &Self) -> bool {
54 self.values.eq(&other.values)
55 }
56}
57
58impl<C: LaunchArg> core::fmt::Debug for SequenceCompilationArg<C> {
59 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
60 f.write_str("Sequence ")?;
61 self.values.fmt(f)
62 }
63}
64impl<C: LaunchArg> core::cmp::Eq for SequenceCompilationArg<C> {}
65
66impl<C: LaunchArg + CubeType + 'static> LaunchArg for Sequence<C> {
67 type RuntimeArg<R: Runtime> = SequenceArg<R, C>;
68 type CompilationArg = SequenceCompilationArg<C>;
69
70 fn register<R: Runtime>(
71 arg: Self::RuntimeArg<R>,
72 launcher: &mut KernelLauncher<R>,
73 ) -> Self::CompilationArg {
74 arg.values
75 .into_iter()
76 .map(|arg| C::register(arg, launcher))
77 .collect()
78 }
79
80 fn expand(arg: &Self::CompilationArg, builder: &mut KernelBuilder) -> SequenceExpand<C> {
81 let values = arg
82 .values
83 .iter()
84 .map(|value| C::expand(value, builder))
85 .collect::<Vec<_>>();
86
87 SequenceExpand { values }
88 }
89}
90
91impl<R: Runtime, E: LaunchArg> FromIterator<E::RuntimeArg<R>> for SequenceArg<R, E> {
92 fn from_iter<T: IntoIterator<Item = E::RuntimeArg<R>>>(iter: T) -> Self {
93 SequenceArg {
94 values: iter.into_iter().collect(),
95 }
96 }
97}
98
99impl<E: LaunchArg> FromIterator<E::CompilationArg> for SequenceCompilationArg<E> {
100 fn from_iter<T: IntoIterator<Item = E::CompilationArg>>(iter: T) -> Self {
101 Self {
102 values: iter.into_iter().collect(),
103 }
104 }
105}
106
107impl<R: Runtime, E: LaunchArg, const N: usize> From<[E::RuntimeArg<R>; N]> for SequenceArg<R, E> {
108 fn from(value: [E::RuntimeArg<R>; N]) -> Self {
109 let mut arg = SequenceArg::<R, E> {
110 values: SmallVec::new(),
111 };
112 for v in value {
113 arg.values.push(v)
114 }
115 arg
116 }
117}