oxidelake_compute/exec/
vector.rs1use std::fmt;
4use std::sync::Arc;
5
6use arrow::datatypes::{DataType, Field, Schema, SchemaRef};
7use datafusion::error::Result;
8use datafusion::execution::TaskContext;
9use datafusion::physical_plan::execution_plan::EmissionType;
10use datafusion::physical_plan::metrics::{BaselineMetrics, ExecutionPlanMetricsSet, MetricsSet};
11use datafusion::physical_plan::stream::RecordBatchStreamAdapter;
12use datafusion::physical_plan::{
13 DisplayAs, DisplayFormatType, ExecutionPlan, ExecutionPlanProperties, PlanProperties,
14 SendableRecordBatchStream,
15};
16use futures::StreamExt;
17use oxidelake_core::BackendKind;
18use oxidelake_core::params::{DistanceMetric, vector_dimension};
19use oxidelake_core::telemetry::TelemetryHub;
20use oxidelake_device::GpuBackend;
21
22use super::{ExecConfig, field_type, plan_err, plan_properties};
23
24#[derive(Debug)]
27pub struct GpuVectorDistanceExec {
28 input: Arc<dyn ExecutionPlan>,
29 column: usize,
30 query: Arc<Vec<f32>>,
31 metric: DistanceMetric,
32 output_name: String,
33 schema: SchemaRef,
34 properties: Arc<PlanProperties>,
35 metrics: ExecutionPlanMetricsSet,
36 config: ExecConfig,
37}
38
39impl GpuVectorDistanceExec {
40 pub fn try_new(
42 input: Arc<dyn ExecutionPlan>,
43 column: usize,
44 query: Vec<f32>,
45 metric: DistanceMetric,
46 output_name: impl Into<String>,
47 target: BackendKind,
48 ) -> Result<Self> {
49 let input_schema = input.schema();
50 let dt = field_type(&input_schema, column, "GpuVectorDistanceExec")?;
51 let dim = vector_dimension(dt)
52 .ok_or_else(|| datafusion::error::DataFusionError::Plan(format!(
53 "GpuVectorDistanceExec: column {column} has type {dt:?}; expected FixedSizeList<Float32>"
54 )))?;
55 if query.len() != dim {
56 return plan_err(format!(
57 "GpuVectorDistanceExec: query has {} dimensions but the column has {dim}",
58 query.len()
59 ));
60 }
61 let output_name = output_name.into();
62 let mut fields: Vec<_> = input_schema.fields().iter().cloned().collect();
63 fields.push(Arc::new(Field::new(&output_name, DataType::Float32, true)));
64 let schema = Arc::new(Schema::new(fields));
65 let properties = plan_properties(
66 Arc::clone(&schema),
67 input.output_partitioning().clone(),
68 EmissionType::Incremental,
69 );
70 Ok(Self {
71 input,
72 column,
73 query: Arc::new(query),
74 metric,
75 output_name,
76 schema,
77 properties,
78 metrics: ExecutionPlanMetricsSet::new(),
79 config: ExecConfig::new(target),
80 })
81 }
82
83 pub fn with_backend(mut self, backend: Arc<dyn GpuBackend>) -> Self {
85 self.config.backend = Some(backend);
86 self
87 }
88
89 pub fn with_telemetry(mut self, telemetry: Arc<TelemetryHub>) -> Self {
91 self.config.telemetry = Some(telemetry);
92 self
93 }
94
95 pub fn column(&self) -> usize {
97 self.column
98 }
99
100 pub fn query(&self) -> &[f32] {
102 &self.query
103 }
104
105 pub fn metric(&self) -> DistanceMetric {
107 self.metric
108 }
109
110 pub fn output_name(&self) -> &str {
112 &self.output_name
113 }
114
115 pub fn target(&self) -> BackendKind {
117 self.config.target
118 }
119
120 pub fn input(&self) -> &Arc<dyn ExecutionPlan> {
122 &self.input
123 }
124}
125
126impl DisplayAs for GpuVectorDistanceExec {
127 fn fmt_as(&self, t: DisplayFormatType, f: &mut fmt::Formatter<'_>) -> fmt::Result {
128 let name = self
129 .input
130 .schema()
131 .fields()
132 .get(self.column)
133 .map_or_else(|| format!("col{}", self.column), |fl| fl.name().clone());
134 match t {
135 DisplayFormatType::Default | DisplayFormatType::Verbose => write!(
136 f,
137 "GpuVectorDistanceExec[{}]: {}({}) AS {}, dim={}",
138 self.config.target,
139 self.metric.name(),
140 name,
141 self.output_name,
142 self.query.len()
143 ),
144 DisplayFormatType::TreeRender => {
145 write!(
146 f,
147 "backend={}\nmetric={}\ncolumn={name}",
148 self.config.target,
149 self.metric.name()
150 )
151 }
152 }
153 }
154}
155
156impl ExecutionPlan for GpuVectorDistanceExec {
157 fn name(&self) -> &str {
158 "GpuVectorDistanceExec"
159 }
160
161 fn properties(&self) -> &Arc<PlanProperties> {
162 &self.properties
163 }
164
165 fn children(&self) -> Vec<&Arc<dyn ExecutionPlan>> {
166 vec![&self.input]
167 }
168
169 fn with_new_children(
170 self: Arc<Self>,
171 children: Vec<Arc<dyn ExecutionPlan>>,
172 ) -> Result<Arc<dyn ExecutionPlan>> {
173 let [input] = <[Arc<dyn ExecutionPlan>; 1]>::try_from(children).map_err(|c| {
174 datafusion::error::DataFusionError::Plan(format!(
175 "GpuVectorDistanceExec expects 1 child, got {}",
176 c.len()
177 ))
178 })?;
179 let mut exec = Self::try_new(
180 input,
181 self.column,
182 self.query.as_ref().clone(),
183 self.metric,
184 self.output_name.clone(),
185 self.config.target,
186 )?;
187 exec.config = self.config.clone();
188 Ok(Arc::new(exec))
189 }
190
191 fn execute(
192 &self,
193 partition: usize,
194 context: Arc<TaskContext>,
195 ) -> Result<SendableRecordBatchStream> {
196 let input = self.input.execute(partition, context)?;
197 let operator = self.config.operator("GpuVectorDistanceExec")?;
198 let baseline = BaselineMetrics::new(&self.metrics, partition);
199 let column = self.column;
200 let query = Arc::clone(&self.query);
201 let metric = self.metric;
202 let output_name = self.output_name.clone();
203 let stream = input.map(move |batch| {
204 let batch = batch?;
205 let _timer = baseline.elapsed_compute().timer();
206 let out = operator.vector_distance(&batch, column, &query, metric, &output_name)?;
207 baseline.record_output(out.num_rows());
208 Ok(out)
209 });
210 Ok(Box::pin(RecordBatchStreamAdapter::new(
211 Arc::clone(&self.schema),
212 stream,
213 )))
214 }
215
216 fn metrics(&self) -> Option<MetricsSet> {
217 Some(self.metrics.clone_inner())
218 }
219}