datafusion_functions_nested/
array_normalize.rs1use crate::utils::make_scalar_function;
21use arrow::array::{
22 Array, ArrayRef, Float64Array, GenericListArray, NullBufferBuilder, OffsetSizeTrait,
23};
24use arrow::buffer::OffsetBuffer;
25use arrow::datatypes::{
26 DataType,
27 DataType::{FixedSizeList, LargeList, List, Null},
28 Field,
29};
30use datafusion_common::cast::{as_float64_array, as_generic_list_array};
31use datafusion_common::utils::{ListCoercion, coerced_type_with_base_type_only};
32use datafusion_common::{Result, internal_err, plan_err, utils::take_function_args};
33use datafusion_expr::{
34 ColumnarValue, Documentation, ScalarFunctionArgs, ScalarUDFImpl, Signature,
35 Volatility,
36};
37use datafusion_macros::user_doc;
38use std::sync::Arc;
39
40make_udf_expr_and_func!(
41 ArrayNormalize,
42 array_normalize,
43 array,
44 "returns the L2-normalized vector for a numeric array.",
45 array_normalize_udf
46);
47
48#[user_doc(
49 doc_section(label = "Array Functions"),
50 description = "Returns the L2-normalized vector for the input numeric array, computed as `array[i] / sqrt(sum(array[i]^2))` per element. Returns NULL if the input is NULL, contains NULL elements, or has zero magnitude (all elements are zero). Returns an empty array for an empty input array.",
51 syntax_example = "array_normalize(array)",
52 sql_example = r#"```sql
53> select array_normalize([3.0, 4.0]);
54+-----------------------------+
55| array_normalize(List([3.0,4.0])) |
56+-----------------------------+
57| [0.6, 0.8] |
58+-----------------------------+
59```"#,
60 argument(
61 name = "array",
62 description = "Array expression. Can be a constant, column, or function, and any combination of array operators."
63 )
64)]
65#[derive(Debug, PartialEq, Eq, Hash)]
66pub struct ArrayNormalize {
67 signature: Signature,
68 aliases: Vec<String>,
69}
70
71impl Default for ArrayNormalize {
72 fn default() -> Self {
73 Self::new()
74 }
75}
76
77impl ArrayNormalize {
78 pub fn new() -> Self {
79 Self {
80 signature: Signature::user_defined(Volatility::Immutable),
81 aliases: vec!["list_normalize".to_string()],
82 }
83 }
84}
85
86impl ScalarUDFImpl for ArrayNormalize {
87 fn name(&self) -> &str {
88 "array_normalize"
89 }
90
91 fn signature(&self) -> &Signature {
92 &self.signature
93 }
94
95 fn return_type(&self, arg_types: &[DataType]) -> Result<DataType> {
96 Ok(arg_types[0].clone())
98 }
99
100 fn coerce_types(&self, arg_types: &[DataType]) -> Result<Vec<DataType>> {
101 let [arg_type] = take_function_args(self.name(), arg_types)?;
102 let coercion = Some(&ListCoercion::FixedSizedListToList);
103
104 if !matches!(arg_type, Null | List(_) | LargeList(_) | FixedSizeList(..)) {
105 return plan_err!("{} does not support type {arg_type}", self.name());
106 }
107
108 let coerced = if matches!(arg_type, Null) {
109 List(Arc::new(Field::new_list_field(DataType::Float64, true)))
110 } else {
111 coerced_type_with_base_type_only(arg_type, &DataType::Float64, coercion)
112 };
113
114 Ok(vec![coerced])
115 }
116
117 fn invoke_with_args(&self, args: ScalarFunctionArgs) -> Result<ColumnarValue> {
118 make_scalar_function(array_normalize_inner)(&args.args)
119 }
120
121 fn aliases(&self) -> &[String] {
122 &self.aliases
123 }
124
125 fn documentation(&self) -> Option<&Documentation> {
126 self.doc()
127 }
128}
129
130fn array_normalize_inner(args: &[ArrayRef]) -> Result<ArrayRef> {
131 let [array] = take_function_args("array_normalize", args)?;
132 match array.data_type() {
133 List(_) => general_array_normalize::<i32>(args),
134 LargeList(_) => general_array_normalize::<i64>(args),
135 arg_type => internal_err!(
136 "array_normalize received unexpected type after coercion: {arg_type}"
137 ),
138 }
139}
140
141fn general_array_normalize<O: OffsetSizeTrait>(arrays: &[ArrayRef]) -> Result<ArrayRef> {
142 let list_array = as_generic_list_array::<O>(&arrays[0])?;
143 let values = as_float64_array(list_array.values())?;
144 let offsets = list_array.value_offsets();
145
146 let mut new_values: Vec<f64> = Vec::with_capacity(values.len());
147 let mut new_offsets = Vec::<O>::with_capacity(list_array.len() + 1);
148 new_offsets.push(O::zero());
149 let mut nulls = NullBufferBuilder::new(list_array.len());
150
151 for row in 0..list_array.len() {
152 if list_array.is_null(row) {
153 nulls.append_null();
154 new_offsets.push(new_offsets[row]);
155 continue;
156 }
157
158 let start = offsets[row].as_usize();
159 let end = offsets[row + 1].as_usize();
160 let len = end - start;
161
162 let slice = values.slice(start, len);
163 if slice.null_count() != 0 {
164 nulls.append_null();
165 new_offsets.push(new_offsets[row]);
166 continue;
167 }
168
169 let vals = slice.values();
170
171 if len == 0 {
173 nulls.append_non_null();
174 new_offsets.push(new_offsets[row]);
175 continue;
176 }
177
178 let mut sq_sum = 0.0;
180 for i in 0..len {
181 sq_sum += vals[i] * vals[i];
182 }
183
184 if sq_sum == 0.0 {
186 nulls.append_null();
187 new_offsets.push(new_offsets[row]);
188 continue;
189 }
190
191 let mag = sq_sum.sqrt();
192 for i in 0..len {
193 new_values.push(vals[i] / mag);
194 }
195 nulls.append_non_null();
196 new_offsets.push(new_offsets[row] + O::usize_as(len));
197 }
198
199 let values_array = Arc::new(Float64Array::from(new_values));
200 let field = Arc::new(Field::new_list_field(DataType::Float64, true));
201
202 Ok(Arc::new(GenericListArray::<O>::try_new(
203 field,
204 OffsetBuffer::new(new_offsets.into()),
205 values_array,
206 nulls.finish(),
207 )?))
208}