// Licensed to the LF AI & Data foundation under one
// or more contributor license agreements. See the NOTICE file
// distributed with this work for additional information
// regarding copyright ownership. The ASF licenses this file
// to you under the Apache License, Version 2.0 (the
// "License"); you may not use this file except in compliance
// with the License. You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
//
// 回读侧的 float16 / bfloat16 解码头。
//
// 写侧目前只有 `@entity.bfloat16_vector_bytes`,回读侧的半边放在这里,因为
// 它只服务于「把一段 payload 变成浮点数组」这一件事;等写侧补齐 float16
// 编码时,两边应当合到 `entity` 去。解码逐位对齐 IEEE-754 binary16,
// 不是近似。
///|
/// 小端两字节一组的 float16 payload 解成浮点数组。
///
/// 奇数字节数是坏负载,直接报错——上游
/// `entity.DeserializeSliceSparseEmbedding` 对不整的字节数也是报错,不截断。
fn float16_bytes_to_floats(
field_name : String,
bytes : Bytes,
) -> Array[Float] raise ColumnError {
if bytes.length() % 2 != 0 {
raise MalformedPayload(
"float16 vector field \{field_name} payload must have an even byte count, got \{bytes.length()}",
)
}
let values : Array[Float] = []
for i = 0; i * 2 < bytes.length(); i = i + 1 {
let low = bytes[i * 2].to_uint()
let high = bytes[i * 2 + 1].to_uint()
values.push(float16_to_float(low | (high << 8)))
}
values
}
///|
/// 一个 IEEE-754 binary16 位模式对应的 float32 值。
///
/// 三种情形分开处理,与 binary16 的定义逐条对应:
/// - 指数全 0:0 或次正规数,尾数要按 2^-14 的标度展开;
/// - 指数全 1:正负无穷或 NaN,保留 NaN 的尾数位(不静默变成 0);
/// - 其余:正规数,指数偏置从 15 换到 127。
fn float16_to_float(bits : UInt) -> Float {
let sign = (bits >> 15) & 1U
let exponent = (bits >> 10) & 0x1FU
let mantissa = bits & 0x3FFU
let sign_bit = sign << 31
if exponent == 0U {
if mantissa == 0U {
return Float::reinterpret_from_uint(sign_bit)
}
// 次正规数:按尾数里最高有效位定位,归一化后是 binary32 的正规数。
let mut shift = 10U
let mut m = mantissa
while (m & 0x400U) == 0U {
m = m << 1
shift = shift - 1U
}
let frac = (m & 0x3FFU) << 13
let exp = (127 - 15 - shift.reinterpret_as_int() + 1).reinterpret_as_uint() <<
23
return Float::reinterpret_from_uint(sign_bit | exp | frac)
}
if exponent == 0x1FU {
// 无穷或 NaN。NaN 必须有非零尾数,binary16 的尾数搬到 binary32 的高位。
if mantissa == 0U {
return Float::reinterpret_from_uint(sign_bit | 0x7F800000U)
}
return Float::reinterpret_from_uint(
sign_bit | 0x7F800000U | (mantissa << 13),
)
}
let exp = (exponent.reinterpret_as_int() - 15 + 127).reinterpret_as_uint() <<
23
Float::reinterpret_from_uint(sign_bit | exp | (mantissa << 13))
}
///|
/// 小端两字节一组的 bfloat16 payload 解成浮点数组。与 `entity` 里的
/// 编码侧互补:那边取 float32 的高 16 位,这边把它放回去,是精确的。
fn bfloat16_bytes_to_floats(
field_name : String,
bytes : Bytes,
) -> Array[Float] raise ColumnError {
if bytes.length() % 2 != 0 {
raise MalformedPayload(
"bfloat16 vector field \{field_name} payload must have an even byte count, got \{bytes.length()}",
)
}
let values : Array[Float] = []
for i = 0; i * 2 < bytes.length(); i = i + 1 {
let low = bytes[i * 2].to_uint()
let high = bytes[i * 2 + 1].to_uint()
values.push(@entity.bfloat16_to_float(low | (high << 8)))
}
values
}