1use polars_core::utils::{
2 _set_partition_size, CustomIterTools, NoNull, accumulate_dataframes_vertical_unchecked,
3 concat_df_unchecked, split,
4};
5use polars_utils::pl_str::PlSmallStr;
6
7use super::*;
8
9fn slice_take(
10 total_rows: IdxSize,
11 n_rows_right: IdxSize,
12 slice: Option<(i64, usize)>,
13 inner: fn(IdxSize, IdxSize, IdxSize) -> IdxCa,
14) -> IdxCa {
15 match slice {
16 None => inner(0, total_rows, n_rows_right),
17 Some((offset, len)) => {
18 let (offset, len) = slice_offsets(offset, len, total_rows as usize);
19 inner(offset as IdxSize, (len + offset) as IdxSize, n_rows_right)
20 },
21 }
22}
23
24fn take_left(total_rows: IdxSize, n_rows_right: IdxSize, slice: Option<(i64, usize)>) -> IdxCa {
25 fn inner(offset: IdxSize, total_rows: IdxSize, n_rows_right: IdxSize) -> IdxCa {
26 let mut take: NoNull<IdxCa> = (offset..total_rows)
27 .map(|i| i / n_rows_right)
28 .collect_trusted();
29 take.set_sorted_flag(IsSorted::Ascending);
30 take.into_inner()
31 }
32 slice_take(total_rows, n_rows_right, slice, inner)
33}
34
35fn take_right(total_rows: IdxSize, n_rows_right: IdxSize, slice: Option<(i64, usize)>) -> IdxCa {
36 fn inner(offset: IdxSize, total_rows: IdxSize, n_rows_right: IdxSize) -> IdxCa {
37 let take: NoNull<IdxCa> = (offset..total_rows)
38 .map(|i| i % n_rows_right)
39 .collect_trusted();
40 take.into_inner()
41 }
42 slice_take(total_rows, n_rows_right, slice, inner)
43}
44
45pub trait CrossJoin: IntoDf {
46 fn cross_join(
48 &self,
49 other: &DataFrame,
50 suffix: Option<PlSmallStr>,
51 slice: Option<(i64, usize)>,
52 maintain_order: MaintainOrderJoin,
53 ) -> PolarsResult<DataFrame> {
54 let (l_df, r_df) = cross_join_dfs(self.to_df(), other, slice, true, maintain_order)?;
55
56 _finish_join(l_df, r_df, suffix)
57 }
58}
59
60impl CrossJoin for DataFrame {}
61
62fn cross_join_dfs<'a>(
63 mut df_self: &'a DataFrame,
64 mut other: &'a DataFrame,
65 slice: Option<(i64, usize)>,
66 parallel: bool,
67 maintain_order: MaintainOrderJoin,
68) -> PolarsResult<(DataFrame, DataFrame)> {
69 if df_self.height() == 0 || other.height() == 0 {
70 return Ok((df_self.clear(), other.clear()));
71 }
72
73 let left_is_primary = match maintain_order {
74 MaintainOrderJoin::None => true,
75 MaintainOrderJoin::Left | MaintainOrderJoin::LeftRight => true,
76 MaintainOrderJoin::Right | MaintainOrderJoin::RightLeft => false,
77 };
78
79 if !left_is_primary {
80 core::mem::swap(&mut df_self, &mut other);
81 }
82
83 let n_rows_left = df_self.height() as IdxSize;
84 let n_rows_right = other.height() as IdxSize;
85 let Some(total_rows) = n_rows_left.checked_mul(n_rows_right) else {
86 polars_bail!(
87 ComputeError: "cross joins would produce more rows than fits into 2^32; \
88 consider compiling with polars-big-idx feature, or set 'streaming'"
89 );
90 };
91
92 let create_left_df = || {
101 unsafe {
104 df_self.take_unchecked_impl(&take_left(total_rows, n_rows_right, slice), parallel)
105 }
106 };
107
108 let create_right_df = || {
109 if n_rows_left > 100 || slice.is_some() {
113 unsafe {
116 other.take_unchecked_impl(&take_right(total_rows, n_rows_right, slice), parallel)
117 }
118 } else {
119 let iter = (0..n_rows_left).map(|_| other);
120 concat_df_unchecked(iter)
121 }
122 };
123 let (l_df, r_df) = if parallel {
124 try_raise_polars_abort();
125 RAYON.install(|| rayon::join(create_left_df, create_right_df))
126 } else {
127 (create_left_df(), create_right_df())
128 };
129 if left_is_primary {
130 Ok((l_df, r_df))
131 } else {
132 Ok((r_df, l_df))
133 }
134}
135
136pub(super) fn fused_cross_filter(
137 left: &DataFrame,
138 right: &DataFrame,
139 suffix: Option<PlSmallStr>,
140 cross_join_options: &CrossJoinOptions,
141 maintain_order: MaintainOrderJoin,
142 emit_unmatched_left: bool,
144 how: &JoinType,
145) -> PolarsResult<DataFrame> {
146 let unfiltered_size = (left.height() as u64).saturating_mul(right.height() as u64);
147 let chunk_size = (unfiltered_size / _set_partition_size() as u64).clamp(1, 100_000);
148 let num_chunks = (unfiltered_size / chunk_size).max(1) as usize;
149
150 let left_is_primary = match maintain_order {
151 MaintainOrderJoin::None => true,
152 MaintainOrderJoin::Left | MaintainOrderJoin::LeftRight => true,
153 MaintainOrderJoin::Right | MaintainOrderJoin::RightLeft => false,
154 };
155 polars_ensure!(
158 !emit_unmatched_left || left_is_primary,
159 InvalidOperation: "'maintain_order={:?}' is not supported for `join_where` with 'how'={}",
160 maintain_order,
161 how
162 );
163
164 let split_chunks;
165 let cartesian_prod = if left_is_primary {
166 split_chunks = split(left, num_chunks);
167 split_chunks.iter().map(|l| (l, right)).collect::<Vec<_>>()
168 } else {
169 split_chunks = split(right, num_chunks);
170 split_chunks.iter().map(|r| (left, r)).collect::<Vec<_>>()
171 };
172
173 let names = _finish_join(left.clear(), right.clear(), suffix)?;
174 let rename_names = names.get_column_names();
175 let rename_names = &rename_names[left.width()..];
176 let len_right = right.height();
177
178 let dfs = RAYON
179 .install(|| {
180 cartesian_prod.par_iter().map(|(left_chunk, right_chunk)| {
181 let (mut joined, right_taken) =
182 cross_join_dfs(left_chunk, right_chunk, None, false, maintain_order)?;
183 let mut right_columns = right_taken.into_columns();
184
185 for (c, name) in right_columns.iter_mut().zip(rename_names) {
186 c.rename((*name).clone());
187 }
188
189 unsafe { joined.hstack_mut_unchecked(&right_columns) };
190
191 if !emit_unmatched_left {
192 cross_join_options.predicate.apply(joined)
193 } else {
194 let mask = cross_join_options.predicate.evaluate(&joined)?;
195
196 let len_left = left_chunk.height();
197 debug_assert_eq!(joined.height(), len_left * len_right);
198
199 let mask_arr = mask.rechunk();
202 let mask_arr = mask_arr.downcast_get(0).unwrap();
203 let match_bits = match mask_arr.validity() {
204 Some(validity) => mask_arr.values() & validity,
205 None => mask_arr.values().clone(),
206 };
207
208 let capacity = match_bits.set_bits() + len_left;
212 let mut left_idx: Vec<IdxSize> = Vec::with_capacity(capacity);
213 let mut right_idx: Vec<NullableIdxSize> = Vec::with_capacity(capacity);
214 for i in 0..len_left {
215 let run = match_bits.clone().sliced(i * len_right, len_right);
216 if run.unset_bits() == len_right {
217 left_idx.push(i as IdxSize);
218 right_idx.push(NullableIdxSize::null());
219 } else {
220 for j in run.true_idx_iter() {
221 left_idx.push(i as IdxSize);
222 right_idx.push(NullableIdxSize::from(j as IdxSize));
223 }
224 }
225 }
226
227 let idx = unsafe { IdxCa::mmap_slice(PlSmallStr::EMPTY, &left_idx) };
228 let mut out = unsafe { left_chunk.take_unchecked(&idx) };
229 let mut right_columns = unsafe {
230 IdxCa::with_nullable_idx(&right_idx, |idx| right.take_unchecked(idx))
231 }
232 .into_columns();
233 for (c, name) in right_columns.iter_mut().zip(rename_names) {
234 c.rename((*name).clone());
235 }
236 unsafe { out.hstack_mut_unchecked(&right_columns) };
237
238 Ok(out)
239 }
240 })
241 })
242 .collect::<PolarsResult<Vec<_>>>()?;
243
244 Ok(accumulate_dataframes_vertical_unchecked(dfs))
245}