#![allow(non_camel_case_types)] #![allow(non_snake_case)] #![allow(clippy::missing_safety_doc)] #![allow(clippy::too_many_arguments)] //! COVER dictionary training. //! //! The public dictionary-builder ABI is declared in `lib/dictBuilder/cover.h`. //! This module mirrors the implementation in `cover.c`; the C translation //! unit is now only an ABI anchor so that the remaining C dictionary builders //! can continue to include that header. use std::cmp::Ordering; use std::collections::HashMap; use std::ffi::c_void; use std::mem::size_of; use std::os::raw::{c_int, c_uint}; use std::ptr; use std::slice; use std::sync::{Arc, Condvar, Mutex, OnceLock}; const ZDICT_DICTSIZE_MIN: usize = 256; const COVER_DEFAULT_SPLITPOINT: f64 = 1.0; const MAP_EMPTY_VALUE: u32 = u32::MAX; const COVER_PRIME4BYTES: u32 = 2_654_435_761; const MAX_ERROR_CODE: usize = 120; const ERROR_GENERIC: usize = 0usize.wrapping_sub(1); const ERROR_MEMORY_ALLOCATION: usize = 0usize.wrapping_sub(64); const ERROR_PARAMETER_OUT_OF_BOUND: usize = 0usize.wrapping_sub(42); const ERROR_SRC_SIZE_WRONG: usize = 0usize.wrapping_sub(72); const ERROR_DST_SIZE_TOO_SMALL: usize = 0usize.wrapping_sub(70); #[repr(C)] #[derive(Clone, Copy, Default)] pub struct ZDICT_params_t { pub compressionLevel: c_int, pub notificationLevel: c_uint, pub dictID: c_uint, } #[repr(C)] #[derive(Clone, Copy, Default)] pub struct ZDICT_cover_params_t { pub k: c_uint, pub d: c_uint, pub steps: c_uint, pub nbThreads: c_uint, pub splitPoint: f64, pub shrinkDict: c_uint, pub shrinkDictMaxRegression: c_uint, pub zParams: ZDICT_params_t, } #[repr(C)] #[derive(Clone, Copy, Default, Debug, PartialEq, Eq)] pub struct COVER_segment_t { pub begin: u32, pub end: u32, pub score: u32, } #[repr(C)] #[derive(Clone, Copy, Default, Debug, PartialEq, Eq)] pub struct COVER_epoch_info_t { pub num: u32, pub size: u32, } #[repr(C)] #[derive(Clone, Copy, Default)] pub struct COVER_dictSelection_t { pub dictContent: *mut u8, pub dictSize: usize, pub totalCompressedSize: usize, } unsafe impl Send for COVER_dictSelection_t {} unsafe impl Sync for COVER_dictSelection_t {} #[repr(C)] #[derive(Clone, Copy)] struct CoverMapPair { key: u32, value: u32, } struct CoverMap { data: Vec, size_log: u32, size_mask: u32, } struct CoverContext { samples: *const u8, samples_sizes: *const usize, offsets: Vec, nb_samples: usize, nb_train_samples: usize, suffix_size: usize, freqs: Vec, dmer_at: Vec, d: usize, } /* Samples are immutable and the scratch arrays are never changed after the * context has been initialized. Each worker owns its frequency copy. */ unsafe impl Send for CoverContext {} unsafe impl Sync for CoverContext {} #[derive(Clone, Copy)] struct BestData { dict: *mut u8, dict_capacity: usize, dict_size: usize, parameters: ZDICT_cover_params_t, compressed_size: usize, live_jobs: usize, } unsafe impl Send for BestData {} struct BestSlot { data: Mutex, cond: Condvar, } impl BestSlot { fn new() -> Self { Self { data: Mutex::new(BestData { dict: ptr::null_mut(), dict_capacity: 0, dict_size: 0, parameters: ZDICT_cover_params_t::default(), compressed_size: usize::MAX, live_jobs: 0, }), cond: Condvar::new(), } } } fn best_slots() -> &'static Mutex>> { static SLOTS: OnceLock>>> = OnceLock::new(); SLOTS.get_or_init(|| Mutex::new(HashMap::new())) } #[inline] fn is_error(value: usize) -> bool { value > 0usize.wrapping_sub(MAX_ERROR_CODE) } #[inline] fn error(code: usize) -> usize { 0usize.wrapping_sub(code) } fn try_vec(len: usize) -> Result, ()> { let mut result = Vec::new(); result.try_reserve_exact(len).map_err(|_| ())?; unsafe { result.set_len(len) }; Ok(result) } #[inline] unsafe fn malloc_bytes(size: usize) -> *mut u8 { unsafe { libc::malloc(size) }.cast::() } #[inline] unsafe fn free_bytes(ptr: *mut u8) { unsafe { libc::free(ptr.cast::()) }; } #[inline] unsafe fn copy_bytes(dst: *mut u8, src: *const u8, size: usize) { unsafe { ptr::copy_nonoverlapping(src, dst, size) }; } #[inline] unsafe fn samples_slice<'a>(samples: *const u8, size: usize) -> &'a [u8] { unsafe { slice::from_raw_parts(samples, size) } } #[inline] fn compare_dmers(ctx: &CoverContext, lhs: u32, rhs: u32) -> Ordering { let lhs = unsafe { samples_slice(ctx.samples.add(lhs as usize), ctx.d) }; let rhs = unsafe { samples_slice(ctx.samples.add(rhs as usize), ctx.d) }; lhs.cmp(rhs) } fn map_init(size: u32) -> Result { let size_log = 32 - size.leading_zeros() + 1; let table_size = 1usize.checked_shl(size_log).ok_or(())?; let mut data = try_vec::(table_size)?; for pair in &mut data { pair.key = 0; pair.value = MAP_EMPTY_VALUE; } Ok(CoverMap { data, size_log, size_mask: table_size as u32 - 1, }) } #[inline] fn map_clear(map: &mut CoverMap) { for pair in &mut map.data { pair.value = MAP_EMPTY_VALUE; } } #[inline] fn map_hash(map: &CoverMap, key: u32) -> u32 { (key.wrapping_mul(COVER_PRIME4BYTES)) >> (32 - map.size_log) } fn map_index(map: &CoverMap, key: u32) -> usize { let hash = map_hash(map, key); let mut index = hash; loop { let pair = &map.data[index as usize]; if pair.value == MAP_EMPTY_VALUE || pair.key == key { return index as usize; } index = (index + 1) & map.size_mask; } } fn map_at(map: &mut CoverMap, key: u32) -> &mut u32 { let index = map_index(map, key); let pair = &mut map.data[index]; if pair.value == MAP_EMPTY_VALUE { pair.key = key; pair.value = 0; } &mut pair.value } fn map_remove(map: &mut CoverMap, key: u32) { let mut index = map_index(map, key); if map.data[index].value == MAP_EMPTY_VALUE { return; } let mut del = index; let mut shift = 1u32; loop { index = (index + 1) & map.size_mask as usize; if map.data[index].value == MAP_EMPTY_VALUE { map.data[del].value = MAP_EMPTY_VALUE; return; } let position = &map.data[index]; let distance = (index as u32).wrapping_sub(map_hash(map, position.key)) & map.size_mask; if distance >= shift { map.data[del] = *position; del = index; shift = 1; } else { shift += 1; } } } fn lower_bound(values: &[usize], value: usize) -> usize { let mut first = 0usize; let mut count = values.len(); while count != 0 { let step = count / 2; let index = first + step; if values[index] < value { first = index + 1; count -= step + 1; } else { count = step; } } first } fn context_init( samples: *const u8, samples_sizes: *const usize, nb_samples: usize, d: usize, split_point: f64, ) -> Result { if nb_samples == 0 || samples_sizes.is_null() { return Err(ERROR_SRC_SIZE_WRONG); } let sizes = unsafe { slice::from_raw_parts(samples_sizes, nb_samples) }; let total_samples_size = sizes .iter() .fold(0usize, |sum, size| sum.wrapping_add(*size)); let nb_train_samples = if split_point < 1.0 { ((nb_samples as f64) * split_point) as usize } else { nb_samples }; let nb_test_samples = if split_point < 1.0 { nb_samples.saturating_sub(nb_train_samples) } else { nb_samples }; let training_samples_size = if split_point < 1.0 { sizes[..nb_train_samples] .iter() .fold(0usize, |sum, size| sum.wrapping_add(*size)) } else { total_samples_size }; let _test_samples_size = if split_point < 1.0 { sizes[nb_train_samples..] .iter() .fold(0usize, |sum, size| sum.wrapping_add(*size)) } else { total_samples_size }; let max_dmer = d.max(size_of::()); let max_samples_size = if size_of::() == 8 { usize::MAX } else { 1usize << 30 }; if total_samples_size < max_dmer || total_samples_size >= max_samples_size || nb_train_samples < 5 || nb_test_samples < 1 || training_samples_size < max_dmer { return Err(ERROR_SRC_SIZE_WRONG); } let suffix_size = training_samples_size - max_dmer + 1; let mut suffix = try_vec::(suffix_size).map_err(|_| ERROR_MEMORY_ALLOCATION)?; let dmer_at = try_vec::(suffix_size).map_err(|_| ERROR_MEMORY_ALLOCATION)?; let mut offsets = try_vec::(nb_samples + 1).map_err(|_| ERROR_MEMORY_ALLOCATION)?; offsets[0] = 0; for (index, size) in sizes.iter().enumerate() { offsets[index + 1] = offsets[index].wrapping_add(*size); } for (index, value) in suffix.iter_mut().enumerate() { *value = index as u32; } let mut context = CoverContext { samples, samples_sizes, offsets, nb_samples, nb_train_samples, suffix_size, freqs: Vec::new(), dmer_at, d, }; suffix.sort_by(|lhs, rhs| compare_dmers(&context, *lhs, *rhs)); let mut group_begin = 0usize; while group_begin < suffix_size { let mut group_end = group_begin + 1; while group_end < suffix_size && compare_dmers(&context, suffix[group_begin], suffix[group_end]) == Ordering::Equal { group_end += 1; } let dmer_id = group_begin as u32; let mut frequency = 0u32; let mut current_offset_index = 0usize; let mut current_sample_end = context.offsets[0]; let group = &suffix[group_begin..group_end]; for (relative_index, &position_raw) in group.iter().enumerate() { let position = position_raw as usize; context.dmer_at[position] = dmer_id; if position < current_sample_end { continue; } frequency = frequency.wrapping_add(1); if relative_index + 1 != group.len() { let search_end = context.nb_samples; let relative = lower_bound(&context.offsets[current_offset_index..search_end], position); let sample_end_index = current_offset_index + relative; current_sample_end = context.offsets[sample_end_index.min(context.nb_samples)]; current_offset_index = (sample_end_index + 1).min(context.nb_samples); } } suffix[group_begin] = frequency; group_begin = group_end; } context.freqs = suffix; Ok(context) } fn select_segment( ctx: &CoverContext, freqs: &mut [u32], active_dmers: &mut CoverMap, begin: u32, end: u32, parameters: ZDICT_cover_params_t, ) -> COVER_segment_t { let dmers_in_k = parameters.k - parameters.d + 1; let mut best = COVER_segment_t::default(); let mut active = COVER_segment_t { begin, end: begin, score: 0, }; map_clear(active_dmers); while active.end < end { let new_dmer = ctx.dmer_at[active.end as usize]; let occurrence = map_at(active_dmers, new_dmer); if *occurrence == 0 { active.score = active.score.wrapping_add(freqs[new_dmer as usize]); } active.end += 1; *occurrence += 1; if active.end - active.begin == dmers_in_k + 1 { let deleted_dmer = ctx.dmer_at[active.begin as usize]; let deleted_occurrence = map_at(active_dmers, deleted_dmer); active.begin += 1; *deleted_occurrence -= 1; if *deleted_occurrence == 0 { map_remove(active_dmers, deleted_dmer); active.score = active.score.wrapping_sub(freqs[deleted_dmer as usize]); } } if active.score > best.score { best = active; } } let mut new_begin = best.end; let mut new_end = best.begin; for position in best.begin..best.end { let frequency = freqs[ctx.dmer_at[position as usize] as usize]; if frequency != 0 { new_begin = new_begin.min(position); new_end = position + 1; } } best.begin = new_begin; best.end = new_end; for position in best.begin..best.end { freqs[ctx.dmer_at[position as usize] as usize] = 0; } best } fn build_dictionary( ctx: &CoverContext, freqs: &mut [u32], active_dmers: &mut CoverMap, dict: *mut u8, dict_capacity: usize, parameters: ZDICT_cover_params_t, ) -> usize { let epochs = COVER_computeEpochs( dict_capacity.min(u32::MAX as usize) as u32, ctx.suffix_size.min(u32::MAX as usize) as u32, parameters.k, 4, ); if epochs.num == 0 || epochs.size == 0 { return dict_capacity; } let max_zero_score_run = 10usize.max((epochs.num as usize >> 3).min(100)); let mut zero_score_run = 0usize; let mut tail = dict_capacity; let mut epoch = 0usize; while tail > 0 { let epoch_begin = (epoch * epochs.size as usize).min(ctx.suffix_size) as u32; let epoch_end = (epoch_begin as usize + epochs.size as usize).min(ctx.suffix_size) as u32; let segment = select_segment(ctx, freqs, active_dmers, epoch_begin, epoch_end, parameters); if segment.score == 0 { zero_score_run += 1; if zero_score_run >= max_zero_score_run { break; } } else { zero_score_run = 0; let segment_size = ((segment.end - segment.begin) as usize) .saturating_add(parameters.d as usize) .saturating_sub(1) .min(tail); if segment_size < parameters.d as usize { break; } tail -= segment_size; unsafe { copy_bytes( dict.add(tail), ctx.samples.add(segment.begin as usize), segment_size, ); } } epoch = (epoch + 1) % epochs.num as usize; } tail } #[inline] fn slot_for(best: *mut c_void) -> Option> { if best.is_null() { return None; } let slots = best_slots() .lock() .unwrap_or_else(|poison| poison.into_inner()); slots.get(&(best as usize)).cloned() } fn best_start_slot(slot: &BestSlot) { let mut data = slot .data .lock() .unwrap_or_else(|poison| poison.into_inner()); data.live_jobs = data.live_jobs.wrapping_add(1); } fn best_wait_slot(slot: &BestSlot) { let mut data = slot .data .lock() .unwrap_or_else(|poison| poison.into_inner()); while data.live_jobs != 0 { data = slot .cond .wait(data) .unwrap_or_else(|poison| poison.into_inner()); } } fn best_finish_slot( slot: &BestSlot, parameters: ZDICT_cover_params_t, selection: COVER_dictSelection_t, ) { let mut data = slot .data .lock() .unwrap_or_else(|poison| poison.into_inner()); data.live_jobs = data.live_jobs.wrapping_sub(1); if selection.totalCompressedSize < data.compressed_size { if data.dict.is_null() || data.dict_capacity < selection.dictSize { let replacement = unsafe { malloc_bytes(selection.dictSize) }; if replacement.is_null() && selection.dictSize != 0 { if !data.dict.is_null() { unsafe { free_bytes(data.dict) }; } data.dict = ptr::null_mut(); data.dict_capacity = 0; data.dict_size = 0; data.compressed_size = ERROR_GENERIC; slot.cond.notify_one(); return; } if !data.dict.is_null() { unsafe { free_bytes(data.dict) }; } data.dict = replacement; data.dict_capacity = selection.dictSize; } if !selection.dictContent.is_null() { unsafe { copy_bytes(data.dict, selection.dictContent, selection.dictSize) }; data.dict_size = selection.dictSize; data.parameters = parameters; data.compressed_size = selection.totalCompressedSize; } } if data.live_jobs == 0 { slot.cond.notify_all(); } } #[derive(Clone, Copy)] struct BestSnapshot { dict: *mut u8, dict_size: usize, parameters: ZDICT_cover_params_t, compressed_size: usize, } unsafe impl Send for BestSnapshot {} fn best_snapshot(slot: &BestSlot) -> BestSnapshot { let data = slot .data .lock() .unwrap_or_else(|poison| poison.into_inner()); BestSnapshot { dict: data.dict, dict_size: data.dict_size, parameters: data.parameters, compressed_size: data.compressed_size, } } type ZSTD_CCtx = c_void; type ZSTD_CDict = c_void; unsafe extern "C" { fn ZDICT_finalizeDictionary( dst_dict_buffer: *mut c_void, max_dict_size: usize, dict_content: *const c_void, dict_content_size: usize, samples_buffer: *const c_void, samples_sizes: *const usize, nb_samples: c_uint, parameters: ZDICT_params_t, ) -> usize; fn ZSTD_compressBound(src_size: usize) -> usize; fn ZSTD_createCCtx() -> *mut ZSTD_CCtx; fn ZSTD_freeCCtx(cctx: *mut ZSTD_CCtx) -> usize; fn ZSTD_createCDict( dict: *const c_void, dict_size: usize, compression_level: c_int, ) -> *mut ZSTD_CDict; fn ZSTD_freeCDict(cdict: *mut ZSTD_CDict) -> usize; fn ZSTD_compress_usingCDict( cctx: *mut ZSTD_CCtx, dst: *mut c_void, dst_capacity: usize, src: *const c_void, src_size: usize, cdict: *const ZSTD_CDict, ) -> usize; } #[no_mangle] pub unsafe extern "C" fn COVER_sum(samples_sizes: *const usize, nb_samples: c_uint) -> usize { if nb_samples == 0 { return 0; } if samples_sizes.is_null() { return 0; } unsafe { slice::from_raw_parts(samples_sizes, nb_samples as usize) .iter() .fold(0usize, |sum, size| sum.wrapping_add(*size)) } } #[no_mangle] pub extern "C" fn COVER_computeEpochs( max_dict_size: u32, nb_dmers: u32, k: u32, passes: u32, ) -> COVER_epoch_info_t { let min_epoch_size = k.saturating_mul(10); let k = k.max(1); let passes = passes.max(1); let mut epochs = COVER_epoch_info_t { num: (max_dict_size / k / passes).max(1), size: 0, }; epochs.size = nb_dmers / epochs.num; if epochs.size < min_epoch_size { epochs.size = min_epoch_size.min(nb_dmers); epochs.num = if epochs.size == 0 { 1 } else { nb_dmers.checked_div(epochs.size).unwrap_or(1).max(1) }; } epochs } #[no_mangle] pub extern "C" fn COVER_warnOnSmallCorpus( max_dict_size: usize, nb_dmers: usize, display_level: c_int, ) { let ratio = nb_dmers as f64 / max_dict_size as f64; if ratio < 10.0 && display_level >= 1 { eprintln!( "WARNING: The maximum dictionary size {} is too large compared to the source size {}! size(source)/size(dictionary) = {}, but it should be >= 10! This may lead to a subpar dictionary! We recommend training on sources at least 10x, and preferably 100x the size of the dictionary! ", max_dict_size, nb_dmers, ratio ); } } #[no_mangle] pub unsafe extern "C" fn ZDICT_trainFromBuffer_cover( dict_buffer: *mut c_void, dict_buffer_capacity: usize, samples_buffer: *const c_void, samples_sizes: *const usize, nb_samples: c_uint, mut parameters: ZDICT_cover_params_t, ) -> usize { parameters.splitPoint = 1.0; if parameters.k == 0 || parameters.d == 0 || parameters.k as usize > dict_buffer_capacity || parameters.d > parameters.k || parameters.splitPoint <= 0.0 || parameters.splitPoint > 1.0 { return error(42); } if nb_samples == 0 { return error(72); } if dict_buffer_capacity < ZDICT_DICTSIZE_MIN { return error(70); } let context = match context_init( samples_buffer.cast(), samples_sizes, nb_samples as usize, parameters.d as usize, parameters.splitPoint, ) { Ok(context) => context, Err(error_code) => return error_code, }; let mut active_dmers = match map_init(parameters.k - parameters.d + 1) { Ok(map) => map, Err(()) => return ERROR_MEMORY_ALLOCATION, }; let mut freqs = match context.freqs.clone().try_reserve(0) { Ok(()) => context.freqs.clone(), Err(_) => return ERROR_MEMORY_ALLOCATION, }; let tail = build_dictionary( &context, &mut freqs, &mut active_dmers, dict_buffer.cast(), dict_buffer_capacity, parameters, ); unsafe { ZDICT_finalizeDictionary( dict_buffer, dict_buffer_capacity, dict_buffer.cast::().add(tail).cast(), dict_buffer_capacity - tail, samples_buffer, samples_sizes, nb_samples, parameters.zParams, ) } } #[no_mangle] pub unsafe extern "C" fn COVER_checkTotalCompressedSize( parameters: ZDICT_cover_params_t, samples_sizes: *const usize, samples: *const u8, offsets: *mut usize, nb_train_samples: usize, nb_samples: usize, dict: *mut u8, dict_buffer_capacity: usize, ) -> usize { if samples_sizes.is_null() || samples.is_null() || offsets.is_null() { return ERROR_GENERIC; } let sizes = unsafe { slice::from_raw_parts(samples_sizes, nb_samples) }; let offset_values = unsafe { slice::from_raw_parts(offsets, nb_samples + 1) }; let start = if parameters.splitPoint < 1.0 { nb_train_samples } else { 0 }; let mut max_sample_size = 0usize; for size in sizes.iter().skip(start) { max_sample_size = max_sample_size.max(*size); } let dst_capacity = unsafe { ZSTD_compressBound(max_sample_size) }; if is_error(dst_capacity) { return dst_capacity; } let dst = unsafe { malloc_bytes(dst_capacity) }; let cctx = unsafe { ZSTD_createCCtx() }; let cdict = unsafe { ZSTD_createCDict( dict.cast(), dict_buffer_capacity, parameters.zParams.compressionLevel, ) }; if dst.is_null() || cctx.is_null() || cdict.is_null() { if !cctx.is_null() { unsafe { ZSTD_freeCCtx(cctx) }; } if !cdict.is_null() { unsafe { ZSTD_freeCDict(cdict) }; } if !dst.is_null() { unsafe { free_bytes(dst) }; } return ERROR_GENERIC; } let mut total = dict_buffer_capacity; for index in start..nb_samples { let result = unsafe { ZSTD_compress_usingCDict( cctx, dst.cast(), dst_capacity, samples.add(offset_values[index]).cast(), sizes[index], cdict, ) }; if is_error(result) { total = result; break; } total = total.wrapping_add(result); } unsafe { ZSTD_freeCCtx(cctx); ZSTD_freeCDict(cdict); free_bytes(dst); } total } #[no_mangle] pub unsafe extern "C" fn COVER_best_init(best: *mut c_void) { if best.is_null() { return; } let slot = Arc::new(BestSlot::new()); let mut slots = best_slots() .lock() .unwrap_or_else(|poison| poison.into_inner()); slots.insert(best as usize, slot); } #[no_mangle] pub unsafe extern "C" fn COVER_best_wait(best: *mut c_void) { if let Some(slot) = slot_for(best) { best_wait_slot(&slot); } } #[no_mangle] pub unsafe extern "C" fn COVER_best_destroy(best: *mut c_void) { if best.is_null() { return; } let slot = { let slots = best_slots() .lock() .unwrap_or_else(|poison| poison.into_inner()); slots.get(&(best as usize)).cloned() }; if let Some(slot) = slot { best_wait_slot(&slot); let mut slots = best_slots() .lock() .unwrap_or_else(|poison| poison.into_inner()); slots.remove(&(best as usize)); let snapshot = best_snapshot(&slot); if !snapshot.dict.is_null() { unsafe { free_bytes(snapshot.dict) }; } } } #[no_mangle] pub unsafe extern "C" fn COVER_best_start(best: *mut c_void) { if let Some(slot) = slot_for(best) { best_start_slot(&slot); } } #[no_mangle] pub unsafe extern "C" fn COVER_best_finish( best: *mut c_void, parameters: ZDICT_cover_params_t, selection: COVER_dictSelection_t, ) { if let Some(slot) = slot_for(best) { best_finish_slot(&slot, parameters, selection); } } #[no_mangle] pub extern "C" fn COVER_dictSelectionError(error_code: usize) -> COVER_dictSelection_t { COVER_dictSelection_t { dictContent: ptr::null_mut(), dictSize: 0, totalCompressedSize: error_code, } } #[no_mangle] pub extern "C" fn COVER_dictSelectionIsError(selection: COVER_dictSelection_t) -> c_uint { (is_error(selection.totalCompressedSize) || selection.dictContent.is_null()) as c_uint } #[no_mangle] pub unsafe extern "C" fn COVER_dictSelectionFree(selection: COVER_dictSelection_t) { if !selection.dictContent.is_null() { unsafe { free_bytes(selection.dictContent) }; } } #[no_mangle] pub unsafe extern "C" fn COVER_selectDict( custom_dict_content: *mut u8, dict_buffer_capacity: usize, mut dict_content_size: usize, samples_buffer: *const u8, samples_sizes: *const usize, nb_finalize_samples: c_uint, nb_check_samples: usize, nb_samples: usize, parameters: ZDICT_cover_params_t, offsets: *mut usize, _total_compressed_size: usize, ) -> COVER_dictSelection_t { let mut total_compressed_size: usize; let custom_dict_content_size = dict_content_size; let custom_dict_end = custom_dict_content.wrapping_add(custom_dict_content_size); let largest_dict_buffer = unsafe { malloc_bytes(dict_buffer_capacity) }; let candidate_dict_buffer = unsafe { malloc_bytes(dict_buffer_capacity) }; if (largest_dict_buffer.is_null() || candidate_dict_buffer.is_null()) && dict_buffer_capacity != 0 { if !largest_dict_buffer.is_null() { unsafe { free_bytes(largest_dict_buffer) }; } if !candidate_dict_buffer.is_null() { unsafe { free_bytes(candidate_dict_buffer) }; } return COVER_dictSelectionError(dict_content_size); } unsafe { copy_bytes(largest_dict_buffer, custom_dict_content, dict_content_size); } dict_content_size = unsafe { ZDICT_finalizeDictionary( largest_dict_buffer.cast(), dict_buffer_capacity, custom_dict_content.cast(), dict_content_size, samples_buffer.cast(), samples_sizes, nb_finalize_samples, parameters.zParams, ) }; if is_error(dict_content_size) { unsafe { free_bytes(largest_dict_buffer); free_bytes(candidate_dict_buffer); } return COVER_dictSelectionError(dict_content_size); } total_compressed_size = unsafe { COVER_checkTotalCompressedSize( parameters, samples_sizes, samples_buffer, offsets, nb_check_samples, nb_samples, largest_dict_buffer, dict_content_size, ) }; if is_error(total_compressed_size) { unsafe { free_bytes(largest_dict_buffer); free_bytes(candidate_dict_buffer); } return COVER_dictSelectionError(total_compressed_size); } if parameters.shrinkDict == 0 { unsafe { free_bytes(candidate_dict_buffer) }; return COVER_dictSelection( largest_dict_buffer, dict_content_size, total_compressed_size, ); } let largest_dict = dict_content_size; let largest_compressed = total_compressed_size; let regression_tolerance = parameters.shrinkDictMaxRegression as f64 / 100.0 + 1.0; dict_content_size = ZDICT_DICTSIZE_MIN; while dict_content_size < largest_dict { unsafe { copy_bytes(candidate_dict_buffer, largest_dict_buffer, largest_dict); let content = custom_dict_end.sub(dict_content_size); dict_content_size = ZDICT_finalizeDictionary( candidate_dict_buffer.cast(), dict_buffer_capacity, content.cast(), dict_content_size, samples_buffer.cast(), samples_sizes, nb_finalize_samples, parameters.zParams, ); } if is_error(dict_content_size) { unsafe { free_bytes(largest_dict_buffer); free_bytes(candidate_dict_buffer); } return COVER_dictSelectionError(dict_content_size); } total_compressed_size = unsafe { COVER_checkTotalCompressedSize( parameters, samples_sizes, samples_buffer, offsets, nb_check_samples, nb_samples, candidate_dict_buffer, dict_content_size, ) }; if is_error(total_compressed_size) { unsafe { free_bytes(largest_dict_buffer); free_bytes(candidate_dict_buffer); } return COVER_dictSelectionError(total_compressed_size); } if (total_compressed_size as f64) <= (largest_compressed as f64) * regression_tolerance { unsafe { free_bytes(largest_dict_buffer) }; return COVER_dictSelection( candidate_dict_buffer, dict_content_size, total_compressed_size, ); } dict_content_size = dict_content_size.saturating_mul(2); } unsafe { free_bytes(candidate_dict_buffer) }; COVER_dictSelection(largest_dict_buffer, largest_dict, largest_compressed) } #[inline] fn COVER_dictSelection( dict_content: *mut u8, dict_size: usize, total_compressed_size: usize, ) -> COVER_dictSelection_t { COVER_dictSelection_t { dictContent: dict_content, dictSize: dict_size, totalCompressedSize: total_compressed_size, } } fn try_parameters( ctx: &CoverContext, slot: &BestSlot, dict_buffer_capacity: usize, parameters: ZDICT_cover_params_t, ) { let mut active_dmers = match map_init(parameters.k - parameters.d + 1) { Ok(map) => map, Err(()) => { best_finish_slot(slot, parameters, COVER_dictSelectionError(ERROR_GENERIC)); return; } }; let dict = unsafe { malloc_bytes(dict_buffer_capacity) }; let mut freqs = match ctx.freqs.clone().try_reserve(0) { Ok(()) => ctx.freqs.clone(), Err(_) => Vec::new(), }; if dict.is_null() || freqs.len() != ctx.freqs.len() { if !dict.is_null() { unsafe { free_bytes(dict) }; } best_finish_slot(slot, parameters, COVER_dictSelectionError(ERROR_GENERIC)); return; } let tail = build_dictionary( ctx, &mut freqs, &mut active_dmers, dict, dict_buffer_capacity, parameters, ); let selection = unsafe { COVER_selectDict( dict.add(tail), dict_buffer_capacity, dict_buffer_capacity - tail, ctx.samples, ctx.samples_sizes, ctx.nb_train_samples as c_uint, ctx.nb_train_samples, ctx.nb_samples, parameters, ctx.offsets.as_ptr().cast_mut(), ERROR_GENERIC, ) }; unsafe { free_bytes(dict) }; best_finish_slot(slot, parameters, selection); unsafe { COVER_dictSelectionFree(selection) }; } #[no_mangle] pub unsafe extern "C" fn ZDICT_optimizeTrainFromBuffer_cover( dict_buffer: *mut c_void, dict_buffer_capacity: usize, samples_buffer: *const c_void, samples_sizes: *const usize, nb_samples: c_uint, parameters: *mut ZDICT_cover_params_t, ) -> usize { let parameters_ref = unsafe { &mut *parameters }; let nb_threads = parameters_ref.nbThreads as usize; let split_point = if parameters_ref.splitPoint <= 0.0 { COVER_DEFAULT_SPLITPOINT } else { parameters_ref.splitPoint }; let k_min_d = if parameters_ref.d == 0 { 6 } else { parameters_ref.d }; let k_max_d = if parameters_ref.d == 0 { 8 } else { parameters_ref.d }; let k_min_k = if parameters_ref.k == 0 { 50 } else { parameters_ref.k }; let k_max_k = if parameters_ref.k == 0 { 2000 } else { parameters_ref.k }; let k_steps = if parameters_ref.steps == 0 { 40 } else { parameters_ref.steps }; if split_point <= 0.0 || split_point > 1.0 || k_min_k < k_max_d || k_max_k < k_min_k || nb_samples == 0 || dict_buffer_capacity < ZDICT_DICTSIZE_MIN { return if nb_samples == 0 { ERROR_SRC_SIZE_WRONG } else if dict_buffer_capacity < ZDICT_DICTSIZE_MIN { ERROR_DST_SIZE_TOO_SMALL } else { ERROR_PARAMETER_OUT_OF_BOUND }; } let k_step_size = ((k_max_k - k_min_k) / k_steps).max(1); let _k_iterations = (1 + (k_max_d - k_min_d) / 2) * (1 + (k_max_k - k_min_k) / k_step_size); let mut token = 0u8; let token_ptr = (&mut token as *mut u8).cast::(); unsafe { COVER_best_init(token_ptr) }; let slot = match slot_for(token_ptr) { Some(slot) => slot, None => return ERROR_MEMORY_ALLOCATION, }; let mut warned = false; for d in (k_min_d..=k_max_d).step_by(2) { let context = match context_init( samples_buffer.cast(), samples_sizes, nb_samples as usize, d as usize, split_point, ) { Ok(context) => context, Err(error_code) => { unsafe { COVER_best_destroy(token_ptr) }; return error_code; } }; if !warned { COVER_warnOnSmallCorpus( dict_buffer_capacity, context.suffix_size, parameters_ref.zParams.notificationLevel as c_int, ); warned = true; } let mut jobs = Vec::new(); let mut k = k_min_k; while k <= k_max_k { let mut job = *parameters_ref; job.k = k; job.d = d; job.splitPoint = split_point; job.steps = k_steps; job.shrinkDict = 0; best_start_slot(&slot); jobs.push(job); k = k.saturating_add(k_step_size); if k == 0 { break; } } if nb_threads > 1 && jobs.len() > 1 { let worker_count = nb_threads.min(jobs.len()); std::thread::scope(|scope| { let next_job = Arc::new(Mutex::new(jobs.into_iter())); let context_ref = &context; for _ in 0..worker_count { let next_job = Arc::clone(&next_job); let slot = Arc::clone(&slot); scope.spawn(move || loop { let job = { let mut jobs = next_job.lock().unwrap_or_else(|poison| poison.into_inner()); jobs.next() }; match job { Some(job) => { try_parameters(context_ref, &slot, dict_buffer_capacity, job) } None => break, } }); } }); } else { for job in jobs { try_parameters(&context, &slot, dict_buffer_capacity, job); } } best_wait_slot(&slot); } let snapshot = best_snapshot(&slot); if is_error(snapshot.compressed_size) { let result = snapshot.compressed_size; unsafe { COVER_best_destroy(token_ptr) }; return result; } unsafe { ptr::copy_nonoverlapping(snapshot.dict, dict_buffer.cast(), snapshot.dict_size); *parameters = snapshot.parameters; COVER_best_destroy(token_ptr); } snapshot.dict_size } #[cfg(test)] mod tests { use super::*; #[test] fn epoch_calculation_matches_cover_rules() { assert_eq!( COVER_computeEpochs(1024, 10_000, 100, 4), COVER_epoch_info_t { num: 2, size: 5000 } ); assert_eq!( COVER_computeEpochs(1024, 100, 100, 4), COVER_epoch_info_t { num: 1, size: 100 } ); } #[test] fn sum_handles_empty_and_multiple_samples() { let sizes = [3usize, 5, 8]; assert_eq!(unsafe { COVER_sum(sizes.as_ptr(), 0) }, 0); assert_eq!( unsafe { COVER_sum(sizes.as_ptr(), sizes.len() as c_uint) }, 16 ); } #[test] fn context_groups_repeated_dmers_by_sample() { let samples = [b'a'; 45]; let sizes = [9usize, 9, 9, 9, 9]; let context = context_init(samples.as_ptr(), sizes.as_ptr(), sizes.len(), 4, 1.0) .expect("valid COVER context"); assert_eq!(context.suffix_size, samples.len() - 8 + 1); assert_eq!(context.freqs.len(), context.suffix_size); assert!(context.freqs.iter().any(|frequency| *frequency >= 5)); } }