use libc::{c_char, c_void};
use lightgbm_sys;
use std;
use std::ffi::CString;
use super::{Error, Result};
pub struct Dataset {
pub(super) handle: lightgbm_sys::DatasetHandle,
}
#[link(name = "c")]
impl Dataset {
fn new(handle: lightgbm_sys::DatasetHandle) -> Self {
Dataset { handle }
}
pub fn from_mat(data: Vec<Vec<f64>>, label: Vec<f32>) -> Result<Self> {
let data_length = data.len();
let feature_length = data[0].len();
let params = CString::new("").unwrap();
let label_str = CString::new("label").unwrap();
let reference = std::ptr::null_mut(); let mut handle = std::ptr::null_mut();
let flat_data = data.into_iter().flatten().collect::<Vec<_>>();
lgbm_call!(lightgbm_sys::LGBM_DatasetCreateFromMat(
flat_data.as_ptr() as *const c_void,
lightgbm_sys::C_API_DTYPE_FLOAT64 as i32,
data_length as i32,
feature_length as i32,
1_i32,
params.as_ptr() as *const c_char,
reference,
&mut handle
))?;
lgbm_call!(lightgbm_sys::LGBM_DatasetSetField(
handle,
label_str.as_ptr() as *const c_char,
label.as_ptr() as *const c_void,
data_length as i32,
lightgbm_sys::C_API_DTYPE_FLOAT32 as i32
))?;
Ok(Dataset::new(handle))
}
pub fn from_file(file_path: &str) -> Result<Self> {
let file_path_str = CString::new(file_path).unwrap();
let params = CString::new("").unwrap();
let mut handle = std::ptr::null_mut();
lgbm_call!(lightgbm_sys::LGBM_DatasetCreateFromFile(
file_path_str.as_ptr() as *const c_char,
params.as_ptr() as *const c_char,
std::ptr::null_mut(),
&mut handle
))?;
Ok(Dataset::new(handle))
}
}
impl Drop for Dataset {
fn drop(&mut self) {
lgbm_call!(lightgbm_sys::LGBM_DatasetFree(self.handle)).unwrap();
}
}
#[cfg(test)]
mod tests {
use super::*;
fn read_train_file() -> Result<Dataset> {
Dataset::from_file(&"lightgbm-sys/lightgbm/examples/binary_classification/binary.train")
}
#[test]
fn read_file() {
assert!(read_train_file().is_ok());
}
#[test]
fn from_mat() {
let data = vec![
vec![1.0, 0.1, 0.2, 0.1],
vec![0.7, 0.4, 0.5, 0.1],
vec![0.9, 0.8, 0.5, 0.1],
vec![0.2, 0.2, 0.8, 0.7],
vec![0.1, 0.7, 1.0, 0.9],
];
let label = vec![0.0, 0.0, 0.0, 1.0, 1.0];
let dataset = Dataset::from_mat(data, label);
assert!(dataset.is_ok());
}
}