use super::crfsuite_format::{write_crfsuite_model, CRFsuiteFeature};
use super::model::CRFModel;
use byteorder::{LittleEndian, ReadBytesExt};
use serde::{Deserialize, Serialize};
use std::fs::File;
use std::io::{BufReader, BufWriter, Read, Write};
use std::path::Path;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CRFFormat {
Native,
CRFsuite,
Auto,
}
const NATIVE_MAGIC: &[u8; 4] = b"UTCF";
const CRFSUITE_MAGIC: &[u8; 4] = b"lCRF";
#[allow(dead_code)]
#[derive(Debug, Clone, Serialize, Deserialize)]
struct NativeHeader {
version: u32,
num_labels: u32,
num_attributes: u32,
num_features: u32,
}
pub struct ModelSaver;
impl ModelSaver {
pub fn new() -> Self {
Self
}
pub fn save<P: AsRef<Path>>(
&self,
model: &CRFModel,
path: P,
format: CRFFormat,
) -> Result<(), String> {
match format {
CRFFormat::Native | CRFFormat::Auto => self.save_native(model, path),
CRFFormat::CRFsuite => self.save_crfsuite(model, path),
}
}
fn save_native<P: AsRef<Path>>(&self, model: &CRFModel, path: P) -> Result<(), String> {
let file = File::create(path).map_err(|e| format!("Failed to create file: {}", e))?;
let mut writer = BufWriter::new(file);
writer
.write_all(NATIVE_MAGIC)
.map_err(|e| format!("Failed to write magic: {}", e))?;
bincode::serialize_into(&mut writer, model)
.map_err(|e| format!("Failed to serialize model: {}", e))?;
writer
.flush()
.map_err(|e| format!("Failed to flush: {}", e))?;
Ok(())
}
fn save_crfsuite<P: AsRef<Path>>(&self, model: &CRFModel, path: P) -> Result<(), String> {
let file = File::create(path).map_err(|e| format!("Failed to create file: {}", e))?;
let mut writer = BufWriter::new(file);
let labels: Vec<String> = (0..model.num_labels)
.filter_map(|i| model.labels.get_label(i as u32).map(|s| s.to_string()))
.collect();
let attributes: Vec<String> = (0..model.num_attributes)
.filter_map(|i| model.attributes.get_attr(i as u32).map(|s| s.to_string()))
.collect();
let mut features: Vec<CRFsuiteFeature> = Vec::new();
for (&(attr_id, label_id), &weight) in model.state_weights_iter() {
if weight.abs() > 1e-10 {
features.push(CRFsuiteFeature {
feat_type: 0,
src: attr_id,
dst: label_id,
weight,
});
}
}
for from_label in 0..model.num_labels {
for to_label in 0..model.num_labels {
let weight = model.get_transition(from_label as u32, to_label as u32);
if weight.abs() > 1e-10 {
features.push(CRFsuiteFeature {
feat_type: 1,
src: from_label as u32,
dst: to_label as u32,
weight,
});
}
}
}
write_crfsuite_model(&mut writer, &labels, &attributes, &features)?;
writer
.flush()
.map_err(|e| format!("Failed to flush: {}", e))?;
Ok(())
}
}
impl Default for ModelSaver {
fn default() -> Self {
Self::new()
}
}
pub struct ModelLoader;
impl ModelLoader {
pub fn new() -> Self {
Self
}
pub fn load<P: AsRef<Path>>(&self, path: P, format: CRFFormat) -> Result<CRFModel, String> {
match format {
CRFFormat::Native => self.load_native(path),
CRFFormat::CRFsuite => self.load_crfsuite(path),
CRFFormat::Auto => self.load_auto(path),
}
}
fn load_auto<P: AsRef<Path>>(&self, path: P) -> Result<CRFModel, String> {
let file = File::open(path.as_ref()).map_err(|e| format!("Failed to open file: {}", e))?;
let mut reader = BufReader::new(file);
let mut magic = [0u8; 4];
reader
.read_exact(&mut magic)
.map_err(|e| format!("Failed to read magic: {}", e))?;
drop(reader);
if &magic == NATIVE_MAGIC {
self.load_native(path)
} else if &magic == CRFSUITE_MAGIC {
self.load_crfsuite(path)
} else {
Err(format!("Unknown file format: {:?}", magic))
}
}
fn load_native<P: AsRef<Path>>(&self, path: P) -> Result<CRFModel, String> {
let file = File::open(path).map_err(|e| format!("Failed to open file: {}", e))?;
let mut reader = BufReader::new(file);
let mut magic = [0u8; 4];
reader
.read_exact(&mut magic)
.map_err(|e| format!("Failed to read magic: {}", e))?;
if &magic != NATIVE_MAGIC {
return Err(format!(
"Invalid magic bytes: expected {:?}, got {:?}",
NATIVE_MAGIC, magic
));
}
let model: CRFModel = bincode::deserialize_from(&mut reader)
.map_err(|e| format!("Failed to deserialize model: {}", e))?;
Ok(model)
}
fn load_crfsuite<P: AsRef<Path>>(&self, path: P) -> Result<CRFModel, String> {
let file = File::open(path).map_err(|e| format!("Failed to open file: {}", e))?;
let mut reader = BufReader::new(file);
let header = self.read_crfsuite_header(&mut reader)?;
let mut model = CRFModel::new();
self.read_crfsuite_strings(
&mut reader,
header.off_labels,
header.num_labels,
&mut model,
true,
)?;
self.read_crfsuite_strings(
&mut reader,
header.off_attrs,
header.num_attrs,
&mut model,
false,
)?;
self.read_crfsuite_features(
&mut reader,
header.off_features,
header.num_features,
&mut model,
)?;
Ok(model)
}
fn read_crfsuite_header(&self, reader: &mut BufReader<File>) -> Result<CRFsuiteHeader, String> {
let mut magic = [0u8; 4];
reader
.read_exact(&mut magic)
.map_err(|e| format!("Failed to read magic: {}", e))?;
if &magic != CRFSUITE_MAGIC {
return Err(format!("Invalid CRFsuite magic: {:?}", magic));
}
let size = reader
.read_u32::<LittleEndian>()
.map_err(|e| format!("Failed to read size: {}", e))?;
let mut type_bytes = [0u8; 4];
reader
.read_exact(&mut type_bytes)
.map_err(|e| format!("Failed to read type: {}", e))?;
let version = reader
.read_u32::<LittleEndian>()
.map_err(|e| format!("Failed to read version: {}", e))?;
let num_features = reader
.read_u32::<LittleEndian>()
.map_err(|e| format!("Failed to read num_features: {}", e))?;
let num_labels = reader
.read_u32::<LittleEndian>()
.map_err(|e| format!("Failed to read num_labels: {}", e))?;
let num_attrs = reader
.read_u32::<LittleEndian>()
.map_err(|e| format!("Failed to read num_attrs: {}", e))?;
let off_features = reader
.read_u32::<LittleEndian>()
.map_err(|e| format!("Failed to read off_features: {}", e))?;
let off_labels = reader
.read_u32::<LittleEndian>()
.map_err(|e| format!("Failed to read off_labels: {}", e))?;
let off_attrs = reader
.read_u32::<LittleEndian>()
.map_err(|e| format!("Failed to read off_attrs: {}", e))?;
let off_label_refs = reader
.read_u32::<LittleEndian>()
.map_err(|e| format!("Failed to read off_label_refs: {}", e))?;
let off_attr_refs = reader
.read_u32::<LittleEndian>()
.map_err(|e| format!("Failed to read off_attr_refs: {}", e))?;
Ok(CRFsuiteHeader {
size,
version,
num_features,
num_labels,
num_attrs,
off_features,
off_labels,
off_attrs,
off_label_refs,
off_attr_refs,
})
}
fn read_crfsuite_strings(
&self,
reader: &mut BufReader<File>,
offset: u32,
count: u32,
model: &mut CRFModel,
is_labels: bool,
) -> Result<(), String> {
use std::io::Seek;
use std::io::SeekFrom;
let cqdb_start = offset;
reader
.seek(SeekFrom::Start(cqdb_start as u64))
.map_err(|e| format!("Failed to seek to CQDB: {}", e))?;
let mut cqdb_magic = [0u8; 4];
reader
.read_exact(&mut cqdb_magic)
.map_err(|e| format!("Failed to read CQDB magic: {}", e))?;
if &cqdb_magic != b"CQDB" {
return Err(format!("Invalid CQDB magic: {:?}", cqdb_magic));
}
let _size = reader
.read_u32::<LittleEndian>()
.map_err(|e| format!("Failed to read CQDB size: {}", e))?;
let _flag = reader
.read_u32::<LittleEndian>()
.map_err(|e| format!("Failed to read CQDB flag: {}", e))?;
let _byteorder = reader
.read_u32::<LittleEndian>()
.map_err(|e| format!("Failed to read CQDB byteorder: {}", e))?;
let bwd_size = reader
.read_u32::<LittleEndian>()
.map_err(|e| format!("Failed to read CQDB bwd_size: {}", e))?;
let bwd_offset = reader
.read_u32::<LittleEndian>()
.map_err(|e| format!("Failed to read CQDB bwd_offset: {}", e))?;
reader
.seek(SeekFrom::Start((cqdb_start + bwd_offset) as u64))
.map_err(|e| format!("Failed to seek to backward table: {}", e))?;
let num_to_read = std::cmp::min(count, bwd_size);
let mut bwd_offsets = Vec::with_capacity(num_to_read as usize);
for _ in 0..num_to_read {
let off = reader
.read_u32::<LittleEndian>()
.map_err(|e| format!("Failed to read backward offset: {}", e))?;
bwd_offsets.push(off);
}
for (id, &entry_offset) in bwd_offsets.iter().enumerate() {
if entry_offset == 0 {
continue; }
reader
.seek(SeekFrom::Start((cqdb_start + entry_offset) as u64))
.map_err(|e| {
format!("Failed to seek to entry at offset {}: {}", entry_offset, e)
})?;
let stored_id = reader
.read_u32::<LittleEndian>()
.map_err(|e| format!("Failed to read entry id: {}", e))?;
if stored_id != id as u32 {
}
let ksize = reader
.read_u32::<LittleEndian>()
.map_err(|e| format!("Failed to read key size: {}", e))?;
let mut key_bytes = vec![0u8; ksize as usize];
reader
.read_exact(&mut key_bytes)
.map_err(|e| format!("Failed to read key bytes: {}", e))?;
if let Some(null_pos) = key_bytes.iter().position(|&b| b == 0) {
key_bytes.truncate(null_pos);
}
let s = String::from_utf8_lossy(&key_bytes).to_string();
if is_labels {
model.labels.get_or_insert(&s);
} else {
model.attributes.get_or_insert(&s);
}
}
model.num_labels = model.labels.len();
model.num_attributes = model.attributes.len();
Ok(())
}
fn read_crfsuite_features(
&self,
reader: &mut BufReader<File>,
offset: u32,
count: u32,
model: &mut CRFModel,
) -> Result<(), String> {
use std::io::Seek;
use std::io::SeekFrom;
reader
.seek(SeekFrom::Start(offset as u64))
.map_err(|e| format!("Failed to seek to features: {}", e))?;
let mut magic = [0u8; 4];
reader
.read_exact(&mut magic)
.map_err(|e| format!("Failed to read FEAT magic: {}", e))?;
if &magic != b"FEAT" {
return Err(format!("Invalid FEAT magic: {:?}", magic));
}
let _chunk_size = reader
.read_u32::<LittleEndian>()
.map_err(|e| format!("Failed to read FEAT chunk size: {}", e))?;
let num_features = reader
.read_u32::<LittleEndian>()
.map_err(|e| format!("Failed to read num features: {}", e))?;
let actual_count = if num_features > 0 {
num_features
} else {
count
};
let num_labels = model.num_labels;
if num_labels > 0 {
model.set_transition(0, 0, 0.0);
}
for _ in 0..actual_count {
let feat_type = reader
.read_u32::<LittleEndian>()
.map_err(|e| format!("Failed to read feature type: {}", e))?;
let source = reader
.read_u32::<LittleEndian>()
.map_err(|e| format!("Failed to read source: {}", e))?;
let target = reader
.read_u32::<LittleEndian>()
.map_err(|e| format!("Failed to read target: {}", e))?;
let weight = reader
.read_f64::<LittleEndian>()
.map_err(|e| format!("Failed to read weight: {}", e))?;
match feat_type {
0 => {
model.set_state_weight(source, target, weight);
}
1 => {
model.set_transition(source, target, weight);
}
_ => {
}
}
}
Ok(())
}
}
impl Default for ModelLoader {
fn default() -> Self {
Self::new()
}
}
#[allow(dead_code)]
#[derive(Debug)]
struct CRFsuiteHeader {
size: u32,
version: u32,
num_features: u32,
num_labels: u32,
num_attrs: u32,
off_features: u32,
off_labels: u32,
off_attrs: u32,
off_label_refs: u32,
off_attr_refs: u32,
}
#[cfg(test)]
mod tests {
use super::*;
fn create_test_model() -> CRFModel {
let mut model = CRFModel::with_labels(vec![
"B-PER".to_string(),
"I-PER".to_string(),
"O".to_string(),
]);
let attr1 = model.attributes.get_or_insert("word=hello");
let attr2 = model.attributes.get_or_insert("word=world");
model.num_attributes = model.attributes.len();
model.set_state_weight(attr1, 0, 1.5);
model.set_state_weight(attr2, 2, -0.5);
model.set_transition(0, 1, 0.8);
model.set_transition(1, 2, 0.3);
model.build_features_list();
model
}
#[test]
fn test_format_detection() {
let native = NATIVE_MAGIC;
let crfsuite = CRFSUITE_MAGIC;
assert_ne!(native, crfsuite);
assert_eq!(native.len(), 4);
assert_eq!(crfsuite.len(), 4);
}
#[test]
fn test_native_roundtrip() {
let model = create_test_model();
let temp_dir = std::env::temp_dir();
let temp_path = temp_dir.join("test_crf_model.bin");
let saver = ModelSaver::new();
saver.save(&model, &temp_path, CRFFormat::Native).unwrap();
let loader = ModelLoader::new();
let loaded = loader.load(&temp_path, CRFFormat::Native).unwrap();
assert_eq!(loaded.num_labels, model.num_labels);
assert_eq!(loaded.num_attributes, model.num_attributes);
let attr1 = loaded.attributes.get("word=hello").unwrap();
assert!((loaded.get_state_weight(attr1, 0) - 1.5).abs() < 1e-10);
std::fs::remove_file(temp_path).ok();
}
#[test]
fn test_auto_format_detection() {
let model = create_test_model();
let temp_dir = std::env::temp_dir();
let temp_path = temp_dir.join("test_crf_model_auto.bin");
let saver = ModelSaver::new();
saver.save(&model, &temp_path, CRFFormat::Native).unwrap();
let loader = ModelLoader::new();
let loaded = loader.load(&temp_path, CRFFormat::Auto).unwrap();
assert_eq!(loaded.num_labels, model.num_labels);
std::fs::remove_file(temp_path).ok();
}
}