perf(engine): share immutable daily factor schemas and numeric buffers
This commit is contained in:
@@ -4389,13 +4389,7 @@ fn normalize_factor_snapshots(
|
|||||||
value,
|
value,
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
let already_normalized = snapshot.extra_factors.iter().all(|(field, value)| {
|
let already_normalized = snapshot.extra_factors.has_normalized_finite_entries();
|
||||||
let trimmed = field.as_str().trim().trim_matches('"').trim_matches('\'');
|
|
||||||
!trimmed.is_empty()
|
|
||||||
&& trimmed == field.as_str()
|
|
||||||
&& trimmed.bytes().all(|byte| !byte.is_ascii_uppercase())
|
|
||||||
&& value.is_finite()
|
|
||||||
});
|
|
||||||
if already_normalized {
|
if already_normalized {
|
||||||
return Ok(snapshot);
|
return Ok(snapshot);
|
||||||
}
|
}
|
||||||
@@ -4465,6 +4459,7 @@ fn normalize_daily_snapshot_bundle(
|
|||||||
|row| row.symbol.as_str(),
|
|row| row.symbol.as_str(),
|
||||||
)?;
|
)?;
|
||||||
sort_rows_by_symbol_if_needed(&mut bundle.market, |row| row.symbol.as_str());
|
sort_rows_by_symbol_if_needed(&mut bundle.market, |row| row.symbol.as_str());
|
||||||
|
NumericFactorMap::share_rows(bundle.factors.iter_mut().map(|row| &mut row.extra_factors));
|
||||||
bundle.factors = normalize_factor_snapshots(bundle.factors)?;
|
bundle.factors = normalize_factor_snapshots(bundle.factors)?;
|
||||||
sort_rows_by_symbol_if_needed(&mut bundle.factors, |row| row.symbol.as_str());
|
sort_rows_by_symbol_if_needed(&mut bundle.factors, |row| row.symbol.as_str());
|
||||||
sort_rows_by_symbol_if_needed(&mut bundle.candidates, |row| row.symbol.as_str());
|
sort_rows_by_symbol_if_needed(&mut bundle.candidates, |row| row.symbol.as_str());
|
||||||
|
|||||||
@@ -8,10 +8,36 @@ use serde::de::{MapAccess, Visitor};
|
|||||||
use serde::ser::SerializeMap;
|
use serde::ser::SerializeMap;
|
||||||
use serde::{Deserialize, Deserializer, Serialize, Serializer};
|
use serde::{Deserialize, Deserializer, Serialize, Serializer};
|
||||||
|
|
||||||
|
mod shared_rows;
|
||||||
|
use shared_rows::{SharedIter, SharedRow};
|
||||||
|
|
||||||
/// Sorted numeric fields stored contiguously, without a tree node per snapshot.
|
/// Sorted numeric fields stored contiguously, without a tree node per snapshot.
|
||||||
#[derive(Clone, Default, PartialEq)]
|
#[derive(Clone)]
|
||||||
pub struct NumericFactorMap {
|
pub struct NumericFactorMap {
|
||||||
entries: Vec<(CompactString, f64)>,
|
storage: Storage,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Clone)]
|
||||||
|
enum Storage {
|
||||||
|
Owned(Vec<(CompactString, f64)>),
|
||||||
|
Shared(SharedRow),
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Default for NumericFactorMap {
|
||||||
|
fn default() -> Self {
|
||||||
|
Self::new()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl PartialEq for NumericFactorMap {
|
||||||
|
fn eq(&self, other: &Self) -> bool {
|
||||||
|
self.len() == other.len() && self.iter().eq(other.iter())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn normalized_name(name: &str) -> bool {
|
||||||
|
let trimmed = name.trim().trim_matches('"').trim_matches('\'');
|
||||||
|
!trimmed.is_empty() && trimmed == name && !name.bytes().any(|byte| byte.is_ascii_uppercase())
|
||||||
}
|
}
|
||||||
|
|
||||||
fn compact_key(key: Cow<'static, str>) -> CompactString {
|
fn compact_key(key: Cow<'static, str>) -> CompactString {
|
||||||
@@ -24,32 +50,70 @@ fn compact_key(key: Cow<'static, str>) -> CompactString {
|
|||||||
impl NumericFactorMap {
|
impl NumericFactorMap {
|
||||||
pub const fn new() -> Self {
|
pub const fn new() -> Self {
|
||||||
Self {
|
Self {
|
||||||
entries: Vec::new(),
|
storage: Storage::Owned(Vec::new()),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn len(&self) -> usize {
|
pub fn len(&self) -> usize {
|
||||||
self.entries.len()
|
match &self.storage {
|
||||||
|
Storage::Owned(entries) => entries.len(),
|
||||||
|
Storage::Shared(row) => row.len(),
|
||||||
|
}
|
||||||
}
|
}
|
||||||
pub fn is_empty(&self) -> bool {
|
pub fn is_empty(&self) -> bool {
|
||||||
self.entries.is_empty()
|
self.len() == 0
|
||||||
}
|
}
|
||||||
pub fn clear(&mut self) {
|
pub fn clear(&mut self) {
|
||||||
self.entries.clear();
|
match &mut self.storage {
|
||||||
|
Storage::Owned(entries) => entries.clear(),
|
||||||
|
Storage::Shared(_) => *self = Self::new(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Shares only immutable values; a subsequent mutation owns a private row.
|
||||||
|
pub fn share_rows<'a>(rows: impl IntoIterator<Item = &'a mut Self>) -> usize {
|
||||||
|
shared_rows::share(rows)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub(crate) fn has_normalized_finite_entries(&self) -> bool {
|
||||||
|
match &self.storage {
|
||||||
|
Storage::Owned(entries) => entries.iter().all(|(name, value)| {
|
||||||
|
normalized_name(name) && value.is_finite()
|
||||||
|
}),
|
||||||
|
Storage::Shared(row) => row.has_normalized_finite_entries(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn owned_entries(&mut self) -> &mut Vec<(CompactString, f64)> {
|
||||||
|
if matches!(self.storage, Storage::Shared(_)) {
|
||||||
|
let entries = self.iter().map(|(key, value)| (key.clone(), *value)).collect();
|
||||||
|
self.storage = Storage::Owned(entries);
|
||||||
|
}
|
||||||
|
match &mut self.storage {
|
||||||
|
Storage::Owned(entries) => entries,
|
||||||
|
Storage::Shared(_) => unreachable!("shared row was materialized"),
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn get(&self, key: &str) -> Option<&f64> {
|
pub fn get(&self, key: &str) -> Option<&f64> {
|
||||||
self.entries
|
match &self.storage {
|
||||||
.binary_search_by(|(name, _)| name.as_str().cmp(key))
|
Storage::Owned(entries) => entries
|
||||||
.ok()
|
.binary_search_by(|(name, _)| name.as_str().cmp(key))
|
||||||
.map(|index| &self.entries[index].1)
|
.ok()
|
||||||
|
.map(|index| &entries[index].1),
|
||||||
|
Storage::Shared(row) => row.get(key),
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn get_mut(&mut self, key: &str) -> Option<&mut f64> {
|
pub fn get_mut(&mut self, key: &str) -> Option<&mut f64> {
|
||||||
self.entries
|
if !self.contains_key(key) {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
let entries = self.owned_entries();
|
||||||
|
entries
|
||||||
.binary_search_by(|(name, _)| name.as_str().cmp(key))
|
.binary_search_by(|(name, _)| name.as_str().cmp(key))
|
||||||
.ok()
|
.ok()
|
||||||
.map(|index| &mut self.entries[index].1)
|
.map(|index| &mut entries[index].1)
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn contains_key(&self, key: &str) -> bool {
|
pub fn contains_key(&self, key: &str) -> bool {
|
||||||
@@ -61,45 +125,51 @@ impl NumericFactorMap {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub fn insert_compact(&mut self, key: CompactString, value: f64) -> Option<f64> {
|
pub fn insert_compact(&mut self, key: CompactString, value: f64) -> Option<f64> {
|
||||||
if self
|
let entries = self.owned_entries();
|
||||||
.entries
|
if entries
|
||||||
.last()
|
.last()
|
||||||
.is_none_or(|(last, _)| last.as_str() < key.as_str())
|
.is_none_or(|(last, _)| last.as_str() < key.as_str())
|
||||||
{
|
{
|
||||||
self.entries.push((key, value));
|
entries.push((key, value));
|
||||||
return None;
|
return None;
|
||||||
}
|
}
|
||||||
match self
|
match entries
|
||||||
.entries
|
|
||||||
.binary_search_by(|(name, _)| name.as_str().cmp(key.as_str()))
|
.binary_search_by(|(name, _)| name.as_str().cmp(key.as_str()))
|
||||||
{
|
{
|
||||||
Ok(index) => Some(std::mem::replace(&mut self.entries[index].1, value)),
|
Ok(index) => Some(std::mem::replace(&mut entries[index].1, value)),
|
||||||
Err(index) => {
|
Err(index) => {
|
||||||
self.entries.insert(index, (key, value));
|
entries.insert(index, (key, value));
|
||||||
None
|
None
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn remove(&mut self, key: &str) -> Option<f64> {
|
pub fn remove(&mut self, key: &str) -> Option<f64> {
|
||||||
self.entries
|
if !self.contains_key(key) {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
let entries = self.owned_entries();
|
||||||
|
entries
|
||||||
.binary_search_by(|(name, _)| name.as_str().cmp(key))
|
.binary_search_by(|(name, _)| name.as_str().cmp(key))
|
||||||
.ok()
|
.ok()
|
||||||
.map(|index| self.entries.remove(index).1)
|
.map(|index| entries.remove(index).1)
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn retain(&mut self, mut keep: impl FnMut(&CompactString, &mut f64) -> bool) {
|
pub fn retain(&mut self, mut keep: impl FnMut(&CompactString, &mut f64) -> bool) {
|
||||||
self.entries.retain_mut(|(key, value)| keep(key, value));
|
self.owned_entries().retain_mut(|(key, value)| keep(key, value));
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn iter(&self) -> Iter<'_> {
|
pub fn iter(&self) -> Iter<'_> {
|
||||||
Iter(self.entries.iter())
|
match &self.storage {
|
||||||
|
Storage::Owned(entries) => Iter(IterStorage::Owned(entries.iter())),
|
||||||
|
Storage::Shared(row) => Iter(IterStorage::Shared(row.iter())),
|
||||||
|
}
|
||||||
}
|
}
|
||||||
pub fn keys(&self) -> impl DoubleEndedIterator<Item = &CompactString> + ExactSizeIterator {
|
pub fn keys(&self) -> impl DoubleEndedIterator<Item = &CompactString> + ExactSizeIterator {
|
||||||
self.entries.iter().map(|(key, _)| key)
|
self.iter().map(|(key, _)| key)
|
||||||
}
|
}
|
||||||
pub fn values(&self) -> impl DoubleEndedIterator<Item = &f64> + ExactSizeIterator {
|
pub fn values(&self) -> impl DoubleEndedIterator<Item = &f64> + ExactSizeIterator {
|
||||||
self.entries.iter().map(|(_, value)| value)
|
self.iter().map(|(_, value)| value)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -116,19 +186,33 @@ impl Index<&str> for NumericFactorMap {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub struct Iter<'a>(std::slice::Iter<'a, (CompactString, f64)>);
|
pub struct Iter<'a>(IterStorage<'a>);
|
||||||
|
|
||||||
|
enum IterStorage<'a> {
|
||||||
|
Owned(std::slice::Iter<'a, (CompactString, f64)>),
|
||||||
|
Shared(SharedIter<'a>),
|
||||||
|
}
|
||||||
impl<'a> Iterator for Iter<'a> {
|
impl<'a> Iterator for Iter<'a> {
|
||||||
type Item = (&'a CompactString, &'a f64);
|
type Item = (&'a CompactString, &'a f64);
|
||||||
fn next(&mut self) -> Option<Self::Item> {
|
fn next(&mut self) -> Option<Self::Item> {
|
||||||
self.0.next().map(|(k, v)| (k, v))
|
match &mut self.0 {
|
||||||
|
IterStorage::Owned(iter) => iter.next().map(|(key, value)| (key, value)),
|
||||||
|
IterStorage::Shared(iter) => iter.next(),
|
||||||
|
}
|
||||||
}
|
}
|
||||||
fn size_hint(&self) -> (usize, Option<usize>) {
|
fn size_hint(&self) -> (usize, Option<usize>) {
|
||||||
self.0.size_hint()
|
match &self.0 {
|
||||||
|
IterStorage::Owned(iter) => iter.size_hint(),
|
||||||
|
IterStorage::Shared(iter) => iter.size_hint(),
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
impl DoubleEndedIterator for Iter<'_> {
|
impl DoubleEndedIterator for Iter<'_> {
|
||||||
fn next_back(&mut self) -> Option<Self::Item> {
|
fn next_back(&mut self) -> Option<Self::Item> {
|
||||||
self.0.next_back().map(|(k, v)| (k, v))
|
match &mut self.0 {
|
||||||
|
IterStorage::Owned(iter) => iter.next_back().map(|(key, value)| (key, value)),
|
||||||
|
IterStorage::Shared(iter) => iter.next_back(),
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
impl ExactSizeIterator for Iter<'_> {}
|
impl ExactSizeIterator for Iter<'_> {}
|
||||||
@@ -143,7 +227,11 @@ impl IntoIterator for NumericFactorMap {
|
|||||||
type Item = (CompactString, f64);
|
type Item = (CompactString, f64);
|
||||||
type IntoIter = std::vec::IntoIter<Self::Item>;
|
type IntoIter = std::vec::IntoIter<Self::Item>;
|
||||||
fn into_iter(self) -> Self::IntoIter {
|
fn into_iter(self) -> Self::IntoIter {
|
||||||
self.entries.into_iter()
|
match self.storage {
|
||||||
|
Storage::Owned(entries) => entries.into_iter(),
|
||||||
|
Storage::Shared(row) => row.iter().map(|(key, value)| (key.clone(), *value))
|
||||||
|
.collect::<Vec<_>>().into_iter(),
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -167,7 +255,7 @@ impl FromIterator<(CompactString, f64)> for NumericFactorMap {
|
|||||||
false
|
false
|
||||||
}
|
}
|
||||||
});
|
});
|
||||||
Self { entries }
|
Self { storage: Storage::Owned(entries) }
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
impl Extend<(Cow<'static, str>, f64)> for NumericFactorMap {
|
impl Extend<(Cow<'static, str>, f64)> for NumericFactorMap {
|
||||||
@@ -177,7 +265,7 @@ impl Extend<(Cow<'static, str>, f64)> for NumericFactorMap {
|
|||||||
}
|
}
|
||||||
impl Extend<(CompactString, f64)> for NumericFactorMap {
|
impl Extend<(CompactString, f64)> for NumericFactorMap {
|
||||||
fn extend<T: IntoIterator<Item = (CompactString, f64)>>(&mut self, iter: T) {
|
fn extend<T: IntoIterator<Item = (CompactString, f64)>>(&mut self, iter: T) {
|
||||||
let mut incoming: Self = iter.into_iter().collect();
|
let incoming: Self = iter.into_iter().collect();
|
||||||
if incoming.is_empty() {
|
if incoming.is_empty() {
|
||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
@@ -185,15 +273,17 @@ impl Extend<(CompactString, f64)> for NumericFactorMap {
|
|||||||
*self = incoming;
|
*self = incoming;
|
||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
if self.entries.last().unwrap().0 < incoming.entries[0].0 {
|
let mut incoming = incoming.into_iter().collect::<Vec<_>>();
|
||||||
self.entries.append(&mut incoming.entries);
|
let entries = self.owned_entries();
|
||||||
|
if entries.last().unwrap().0 < incoming[0].0 {
|
||||||
|
entries.append(&mut incoming);
|
||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
// Merge sorted sets in linear time; wide factor batches must not shift
|
// Merge sorted sets in linear time; wide factor batches must not shift
|
||||||
// the existing vector once per field. Existing keys keep their identity.
|
// the existing vector once per field. Existing keys keep their identity.
|
||||||
let mut merged = Vec::with_capacity(self.len() + incoming.len());
|
let mut merged = Vec::with_capacity(entries.len() + incoming.len());
|
||||||
let mut old = std::mem::take(&mut self.entries).into_iter().peekable();
|
let mut old = std::mem::take(entries).into_iter().peekable();
|
||||||
let mut new = incoming.entries.into_iter().peekable();
|
let mut new = incoming.into_iter().peekable();
|
||||||
while let (Some(left), Some(right)) = (old.peek(), new.peek()) {
|
while let (Some(left), Some(right)) = (old.peek(), new.peek()) {
|
||||||
match left.0.cmp(&right.0) {
|
match left.0.cmp(&right.0) {
|
||||||
std::cmp::Ordering::Less => merged.push(old.next().unwrap()),
|
std::cmp::Ordering::Less => merged.push(old.next().unwrap()),
|
||||||
@@ -206,7 +296,7 @@ impl Extend<(CompactString, f64)> for NumericFactorMap {
|
|||||||
}
|
}
|
||||||
merged.extend(old);
|
merged.extend(old);
|
||||||
merged.extend(new);
|
merged.extend(new);
|
||||||
self.entries = merged;
|
*entries = merged;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
impl<const N: usize> From<[(Cow<'static, str>, f64); N]> for NumericFactorMap {
|
impl<const N: usize> From<[(Cow<'static, str>, f64); N]> for NumericFactorMap {
|
||||||
@@ -253,6 +343,99 @@ impl<'de> Deserialize<'de> for NumericFactorMap {
|
|||||||
mod tests {
|
mod tests {
|
||||||
use super::*;
|
use super::*;
|
||||||
|
|
||||||
|
fn row_bits(row: &NumericFactorMap) -> Vec<(String, u64)> {
|
||||||
|
row.iter().map(|(key, value)| (key.to_string(), value.to_bits())).collect()
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn shared_rows_keep_missing_entries_names_and_float_bits() {
|
||||||
|
let mut rows = (0..128).map(|row| {
|
||||||
|
(0..12).filter(|field| (field + row) % 7 != 0).map(|field| {
|
||||||
|
let value = match (row, field) {
|
||||||
|
(1, 1) => -0.0,
|
||||||
|
(2, 1) => f64::from_bits(0x7ff8_0000_0000_0012),
|
||||||
|
(3, 1) => f64::INFINITY,
|
||||||
|
_ => row as f64 / 7.0 + field as f64,
|
||||||
|
};
|
||||||
|
(Cow::Owned(format!("factor_{field:02}")), value)
|
||||||
|
}).collect::<NumericFactorMap>()
|
||||||
|
}).collect::<Vec<_>>();
|
||||||
|
rows[7].clear();
|
||||||
|
rows[8].insert(Cow::Borrowed(" 'Quoted' "), 1.0);
|
||||||
|
let before = rows.iter().map(row_bits).collect::<Vec<_>>();
|
||||||
|
let serialized = rows.iter().map(|row| serde_json::to_string(row).unwrap()).collect::<Vec<_>>();
|
||||||
|
let normalized = rows.iter().map(NumericFactorMap::has_normalized_finite_entries).collect::<Vec<_>>();
|
||||||
|
assert_eq!(NumericFactorMap::share_rows(rows.iter_mut()), rows.len());
|
||||||
|
for (index, row) in rows.iter().enumerate() {
|
||||||
|
assert!(matches!(row.storage, Storage::Shared(_)));
|
||||||
|
assert_eq!(row_bits(row), before[index]);
|
||||||
|
assert_eq!(serde_json::to_string(row).unwrap(), serialized[index]);
|
||||||
|
assert_eq!(row.has_normalized_finite_entries(), normalized[index]);
|
||||||
|
for field in 0..12 {
|
||||||
|
let name = format!("factor_{field:02}");
|
||||||
|
assert_eq!(row.get(&name).map(|value| value.to_bits()),
|
||||||
|
before[index].iter().find(|(key, _)| key == &name).map(|(_, value)| *value));
|
||||||
|
}
|
||||||
|
let mut remaining = before[index].clone();
|
||||||
|
let mut iter = row.iter();
|
||||||
|
while !remaining.is_empty() {
|
||||||
|
assert_eq!(iter.len(), remaining.len());
|
||||||
|
let (expected, actual) = if remaining.len() % 2 == 0 {
|
||||||
|
(remaining.remove(0), iter.next())
|
||||||
|
} else {
|
||||||
|
(remaining.pop().unwrap(), iter.next_back())
|
||||||
|
};
|
||||||
|
let (key, value) = actual.unwrap();
|
||||||
|
assert_eq!((key.to_string(), value.to_bits()), expected);
|
||||||
|
}
|
||||||
|
assert_eq!(iter.len(), 0);
|
||||||
|
assert!(iter.next().is_none() && iter.next_back().is_none());
|
||||||
|
assert_eq!(row.clone().into_iter().map(|(key, value)| (key.to_string(), value.to_bits()))
|
||||||
|
.collect::<Vec<_>>(), before[index]);
|
||||||
|
}
|
||||||
|
assert_eq!(NumericFactorMap::share_rows(rows.iter_mut()), 0);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn shared_row_mutations_are_private_and_keep_map_semantics() {
|
||||||
|
let mut rows = (0..64).map(|row| {
|
||||||
|
(0..8).map(|field| (Cow::Owned(format!("f{field}")), (row * 8 + field) as f64))
|
||||||
|
.collect::<NumericFactorMap>()
|
||||||
|
}).collect::<Vec<_>>();
|
||||||
|
let originals = rows.clone();
|
||||||
|
assert_eq!(NumericFactorMap::share_rows(rows.iter_mut()), 64);
|
||||||
|
assert_eq!(rows, originals);
|
||||||
|
let surviving = rows[0].clone();
|
||||||
|
*rows[0].get_mut("f0").unwrap() = -99.0;
|
||||||
|
rows[1].insert(Cow::Borrowed("f1"), 101.0);
|
||||||
|
rows[2].remove("f2");
|
||||||
|
rows[3].retain(|key, value| { *value *= 2.0; key.as_str() != "f3" });
|
||||||
|
rows[4].extend([(Cow::Borrowed("f1"), 303.0), (Cow::Borrowed("new"), -0.0)]);
|
||||||
|
rows[5].clear();
|
||||||
|
assert_eq!(surviving, originals[0]);
|
||||||
|
assert_eq!(rows[0]["f0"], -99.0);
|
||||||
|
assert_eq!(rows[1]["f1"], 101.0);
|
||||||
|
assert!(!rows[2].contains_key("f2"));
|
||||||
|
assert!(!rows[3].contains_key("f3"));
|
||||||
|
assert_eq!(rows[3]["f0"], originals[3]["f0"] * 2.0);
|
||||||
|
assert_eq!(rows[4]["new"].to_bits(), (-0.0_f64).to_bits());
|
||||||
|
assert!(rows[5].is_empty());
|
||||||
|
assert_eq!(&rows[6..], &originals[6..]);
|
||||||
|
drop(rows);
|
||||||
|
assert_eq!(surviving, originals[0]);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn sparse_rows_do_not_expand_into_a_dense_matrix() {
|
||||||
|
let mut rows = (0..256).map(|row| NumericFactorMap::from([
|
||||||
|
(Cow::Owned(format!("only_{row}")), row as f64),
|
||||||
|
])).collect::<Vec<_>>();
|
||||||
|
let before = rows.clone();
|
||||||
|
assert_eq!(NumericFactorMap::share_rows(rows.iter_mut()), 0);
|
||||||
|
assert_eq!(rows, before);
|
||||||
|
assert!(rows.iter().all(|row| matches!(row.storage, Storage::Owned(_))));
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn compact_keys_inline_dynamic_names_and_keep_long_static_storage() {
|
fn compact_keys_inline_dynamic_names_and_keep_long_static_storage() {
|
||||||
const LONG: &str = "a_long_static_factor_identifier_that_must_remain_borrowed";
|
const LONG: &str = "a_long_static_factor_identifier_that_must_remain_borrowed";
|
||||||
|
|||||||
@@ -0,0 +1,165 @@
|
|||||||
|
use std::collections::BTreeSet;
|
||||||
|
use std::sync::Arc;
|
||||||
|
|
||||||
|
use compact_str::CompactString;
|
||||||
|
|
||||||
|
use super::{NumericFactorMap, Storage, normalized_name};
|
||||||
|
|
||||||
|
const MAX_SHARED_ROWS_BYTES: usize = 64 * 1024 * 1024;
|
||||||
|
|
||||||
|
struct SharedRows {
|
||||||
|
fields: Vec<CompactString>,
|
||||||
|
values: Box<[f64]>,
|
||||||
|
present: Box<[u64]>,
|
||||||
|
lengths: Box<[usize]>,
|
||||||
|
normalized_finite: Box<[bool]>,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Clone)]
|
||||||
|
pub(super) struct SharedRow {
|
||||||
|
data: Arc<SharedRows>,
|
||||||
|
index: usize,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl SharedRow {
|
||||||
|
pub(super) fn len(&self) -> usize {
|
||||||
|
self.data.lengths[self.index]
|
||||||
|
}
|
||||||
|
|
||||||
|
pub(super) fn has_normalized_finite_entries(&self) -> bool {
|
||||||
|
self.data.normalized_finite[self.index]
|
||||||
|
}
|
||||||
|
|
||||||
|
fn value_at(&self, field: usize) -> Option<&f64> {
|
||||||
|
let offset = self.index * self.data.fields.len() + field;
|
||||||
|
(self.data.present[offset / 64] & (1_u64 << (offset % 64)) != 0)
|
||||||
|
.then(|| &self.data.values[offset])
|
||||||
|
}
|
||||||
|
|
||||||
|
pub(super) fn get(&self, name: &str) -> Option<&f64> {
|
||||||
|
self.data.fields.binary_search_by(|field| field.as_str().cmp(name))
|
||||||
|
.ok().and_then(|field| self.value_at(field))
|
||||||
|
}
|
||||||
|
|
||||||
|
pub(super) fn iter(&self) -> SharedIter<'_> {
|
||||||
|
SharedIter { row: self, front: 0, back: self.data.fields.len(), remaining: self.len() }
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub(super) struct SharedIter<'a> {
|
||||||
|
row: &'a SharedRow,
|
||||||
|
front: usize,
|
||||||
|
back: usize,
|
||||||
|
remaining: usize,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl<'a> Iterator for SharedIter<'a> {
|
||||||
|
type Item = (&'a CompactString, &'a f64);
|
||||||
|
|
||||||
|
fn next(&mut self) -> Option<Self::Item> {
|
||||||
|
if self.remaining == 0 {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
while self.front < self.back {
|
||||||
|
let field = self.front;
|
||||||
|
self.front += 1;
|
||||||
|
if let Some(value) = self.row.value_at(field) {
|
||||||
|
self.remaining -= 1;
|
||||||
|
return Some((&self.row.data.fields[field], value));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
None
|
||||||
|
}
|
||||||
|
|
||||||
|
fn size_hint(&self) -> (usize, Option<usize>) {
|
||||||
|
(self.remaining, Some(self.remaining))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl DoubleEndedIterator for SharedIter<'_> {
|
||||||
|
fn next_back(&mut self) -> Option<Self::Item> {
|
||||||
|
if self.remaining == 0 {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
while self.front < self.back {
|
||||||
|
self.back -= 1;
|
||||||
|
if let Some(value) = self.row.value_at(self.back) {
|
||||||
|
self.remaining -= 1;
|
||||||
|
return Some((&self.row.data.fields[self.back], value));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
None
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl ExactSizeIterator for SharedIter<'_> {}
|
||||||
|
|
||||||
|
pub(super) fn share<'a>(rows: impl IntoIterator<Item = &'a mut NumericFactorMap>) -> usize {
|
||||||
|
let rows = rows.into_iter().collect::<Vec<_>>();
|
||||||
|
if rows.len() < 2 || rows.iter().any(|row| matches!(row.storage, Storage::Shared(_))) {
|
||||||
|
return 0;
|
||||||
|
}
|
||||||
|
let Some(owned_bytes) = rows.iter().try_fold(0usize, |sum, row| {
|
||||||
|
let Storage::Owned(entries) = &row.storage else { return None };
|
||||||
|
sum.checked_add(entries.capacity().checked_mul(std::mem::size_of::<(CompactString, f64)>())?)
|
||||||
|
}) else { return 0 };
|
||||||
|
let fields = rows.iter().flat_map(|row| row.keys()).collect::<BTreeSet<_>>()
|
||||||
|
.into_iter().cloned().collect::<Vec<_>>();
|
||||||
|
let Some(cells) = rows.len().checked_mul(fields.len()) else { return 0 };
|
||||||
|
if cells == 0 {
|
||||||
|
return 0;
|
||||||
|
}
|
||||||
|
let words = cells.div_ceil(64);
|
||||||
|
let Some(field_bytes) = fields.iter().try_fold(0usize, |size, field| {
|
||||||
|
size.checked_add(std::mem::size_of::<CompactString>())?.checked_add(field.len())
|
||||||
|
}) else { return 0 };
|
||||||
|
let Some(bytes) = cells.checked_mul(std::mem::size_of::<f64>())
|
||||||
|
.and_then(|size| size.checked_add(words.checked_mul(std::mem::size_of::<u64>())?))
|
||||||
|
.and_then(|size| size.checked_add(rows.len().checked_mul(std::mem::size_of::<usize>() + 1)?))
|
||||||
|
.and_then(|size| size.checked_add(field_bytes))
|
||||||
|
.and_then(|size| size.checked_add(std::mem::size_of::<SharedRows>()))
|
||||||
|
else { return 0 };
|
||||||
|
if bytes >= owned_bytes || bytes > MAX_SHARED_ROWS_BYTES {
|
||||||
|
return 0;
|
||||||
|
}
|
||||||
|
|
||||||
|
let mut values = Vec::new();
|
||||||
|
if values.try_reserve_exact(cells).is_err() {
|
||||||
|
return 0;
|
||||||
|
}
|
||||||
|
values.resize(cells, 0.0);
|
||||||
|
let mut present = vec![0_u64; words];
|
||||||
|
let mut lengths = Vec::with_capacity(rows.len());
|
||||||
|
let mut normalized_finite = Vec::with_capacity(rows.len());
|
||||||
|
let canonical_fields = fields.iter().map(|name| normalized_name(name)).collect::<Vec<_>>();
|
||||||
|
for (index, row) in rows.iter().enumerate() {
|
||||||
|
lengths.push(row.len());
|
||||||
|
let mut canonical_row = true;
|
||||||
|
let mut field = 0usize;
|
||||||
|
for (name, value) in row.iter() {
|
||||||
|
while &fields[field] < name {
|
||||||
|
field += 1;
|
||||||
|
}
|
||||||
|
debug_assert_eq!(&fields[field], name);
|
||||||
|
let offset = index * fields.len() + field;
|
||||||
|
values[offset] = *value;
|
||||||
|
present[offset / 64] |= 1_u64 << (offset % 64);
|
||||||
|
canonical_row &= canonical_fields[field] && value.is_finite();
|
||||||
|
field += 1;
|
||||||
|
}
|
||||||
|
normalized_finite.push(canonical_row);
|
||||||
|
}
|
||||||
|
// Publish only after every input row has been copied without arithmetic.
|
||||||
|
let data = Arc::new(SharedRows {
|
||||||
|
fields,
|
||||||
|
values: values.into_boxed_slice(),
|
||||||
|
present: present.into_boxed_slice(),
|
||||||
|
lengths: lengths.into_boxed_slice(),
|
||||||
|
normalized_finite: normalized_finite.into_boxed_slice(),
|
||||||
|
});
|
||||||
|
let count = rows.len();
|
||||||
|
for (index, row) in rows.into_iter().enumerate() {
|
||||||
|
row.storage = Storage::Shared(SharedRow { data: Arc::clone(&data), index });
|
||||||
|
}
|
||||||
|
count
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user