use polars_async::primitives::distributor_channel::distributor_channel;
use polars_async::primitives::wait_group::WaitGroup;
use polars_core::prelude::{AnyValue, Column, DataType, IntoColumn};
use polars_core::scalar::Scalar;
use polars_error::PolarsResult;
use polars_ops::series::{InterpolationMethod, interpolate};
use polars_utils::IdxSize;
use polars_utils::pl_str::PlSmallStr;
use super::compute_node_prelude::*;
use crate::DEFAULT_DISTRIBUTOR_BUFFER_SIZE;
use crate::morsel::{MorselSeq, SourceToken, get_ideal_morsel_size};
pub struct InterpolateNode {
method: InterpolationMethod,
input_dtype: DataType,
output_dtype: DataType,
col_name: PlSmallStr,
seq: MorselSeq,
last_non_null: AnyValue<'static>,
pending_nulls: IdxSize,
}
impl InterpolateNode {
pub fn new(
method: InterpolationMethod,
input_dtype: DataType,
output_dtype: DataType,
col_name: PlSmallStr,
) -> Self {
Self {
method,
input_dtype,
output_dtype,
col_name,
seq: MorselSeq::default(),
last_non_null: AnyValue::Null,
pending_nulls: 0,
}
}
}
impl ComputeNode for InterpolateNode {
fn name(&self) -> &str {
"interpolate"
}
fn update_state(
&mut self,
recv: &mut [PortState],
send: &mut [PortState],
_state: &StreamingExecutionState,
) -> PolarsResult<()> {
assert!(recv.len() == 1 && send.len() == 1);
if send[0] == PortState::Done {
recv[0] = PortState::Done;
self.pending_nulls = 0;
} else if recv[0] == PortState::Done {
if self.pending_nulls > 0 {
send[0] = PortState::Ready;
} else {
send[0] = PortState::Done;
}
} else {
recv.swap_with_slice(send);
}
Ok(())
}
fn spawn<'env, 's>(
&'env mut self,
scope: &'s TaskScope<'s, 'env>,
recv_ports: &mut [Option<RecvPort<'_>>],
send_ports: &mut [Option<SendPort<'_>>],
_state: &'s StreamingExecutionState,
join_handles: &mut Vec<JoinHandle<PolarsResult<()>>>,
) {
assert_eq!(recv_ports.len(), 1);
assert_eq!(send_ports.len(), 1);
let recv = recv_ports[0].take();
let send = send_ports[0].take().unwrap();
let method = self.method;
let input_dtype = self.input_dtype.clone();
let output_dtype = self.output_dtype.clone();
let col_name = self.col_name.clone();
let pending_nulls = &mut self.pending_nulls;
let last_non_null = &mut self.last_non_null;
let seq = &mut self.seq;
let Some(recv) = recv else {
debug_assert!(*pending_nulls > 0);
let source_token = SourceToken::new();
let mut send = send.serial();
join_handles.push(scope.spawn_task(TaskPriority::High, async move {
let morsel_size = get_ideal_morsel_size();
while *pending_nulls > 0 && !source_token.stop_requested() {
let chunk_size = morsel_size.min(*pending_nulls as usize);
let df =
Column::full_null(col_name.clone(), chunk_size, &output_dtype).into_frame();
if send
.send(Morsel::new(df, *seq, source_token.clone()))
.await
.is_err()
{
break;
}
*seq = seq.successor();
*pending_nulls -= chunk_size as IdxSize;
}
Ok(())
}));
return;
};
let mut receiver = recv.serial();
let senders = send.parallel();
let (mut distributor, distr_receivers) =
distributor_channel(senders.len(), *DEFAULT_DISTRIBUTOR_BUFFER_SIZE);
join_handles.push(scope.spawn_task(TaskPriority::High, async move {
while let Ok(morsel) = receiver.recv().await {
let (df, _, source_token, _) = morsel.into_inner();
let mut columns = df.into_columns();
assert_eq!(columns.len(), 1);
let column = columns.pop().unwrap();
let height = column.len();
let Some(last_non_null_idx) = column.last_non_null() else {
*pending_nulls += height as IdxSize;
continue;
};
let ready_values = column.slice(0, last_non_null_idx + 1);
if distributor
.send((
*seq,
source_token,
*pending_nulls,
last_non_null.clone(),
ready_values,
))
.await
.is_err()
{
return Ok(());
}
*seq = seq.successor();
*last_non_null = column.get(last_non_null_idx).unwrap().into_static();
*pending_nulls = (height - 1 - last_non_null_idx) as IdxSize;
}
Ok(())
}));
for (mut send, mut recv) in senders.into_iter().zip(distr_receivers) {
let input_dtype = input_dtype.clone();
let output_dtype = output_dtype.clone();
join_handles.push(scope.spawn_task(TaskPriority::High, async move {
let wait_group = WaitGroup::default();
while let Ok((seq, source_token, pending_nulls, last_non_null, mut column)) =
recv.recv().await
{
let has_prepended = pending_nulls > 0 || !last_non_null.is_null();
if has_prepended {
let mut c = Column::new_scalar(
column.name().clone(),
Scalar::new(input_dtype.clone(), last_non_null),
1,
);
c.append_owned(Column::full_null(
column.name().clone(),
pending_nulls as usize,
&input_dtype,
))?;
c.append_owned(column)?;
column = c;
}
column = if column.has_nulls() {
interpolate(column.as_materialized_series(), method).into_column()
} else {
column.cast(&output_dtype)?
};
if has_prepended {
column = column.slice(1, usize::MAX);
}
let mut morsel = Morsel::new(column.into_frame(), seq, source_token.clone());
morsel.set_consume_token(wait_group.token());
if send.send(morsel).await.is_err() {
break;
}
wait_group.wait().await;
}
Ok(())
}));
}
}
}