Burn 自定义 CSV 数据集实战:InMemDataset 与 Polars DataframeDataset 两条实现路径

发布时间:2026/9/14 18:18:52
Burn 自定义 CSV 数据集实战:InMemDataset 与 Polars DataframeDataset 两条实现路径 Burn 自定义 CSV 数据集实战InMemDataset 与 Polars DataframeDataset 两条实现路径【免费下载链接】burnBurn is a next generation tensor library and Deep Learning Framework that doesnt compromise on flexibility, efficiency and portability.项目地址: https://gitcode.com/GitHub_Trending/bu/burn本文基于 Burn 仓库中 custom-csv-dataset 示例讲解如何为任意 CSV 文件实现 Burn 的Datasettrait。示例以经典的 diabetes 数据集442 条患者记录为载体给出了两条可落地的技术路线一是用csvserde将文件解析进内存、由InMemDataset承载二是用 Polars 做惰性、列式加载。读完本篇你可以掌握如何定义带serde字段映射的记录结构体、如何为自定义数据集实现Datasettrait、两条路线的适用场景差异以及示例中cargo run --example的运行方式与 feature 开关。示例概览一个数据集两种Dataset实现该示例位于仓库的examples/custom-csv-dataset/目录结构如下src/dataset.rs —— 基于InMemDataset的内存版实现src/dataframe_dataset.rs —— 基于 Polars DataFrame 的DataframeDataset实现src/diabetes_patient.rs —— 单条患者记录DiabetesPatient的serde结构体定义src/utils.rs —— CSV 文件的按需下载逻辑examples/custom-csv-dataset.rs 与 examples/dataframe-dataset.rs —— 两个可运行的入口。数据集本身是 scikit-learn 使用的 diabetes 数据集原始来源为 NCSU 的diabetes.tab.txt包含442 条患者记录每条记录有 10 个基线变量年龄、性别、BMI、平均血压以及 6 项血清指标 S1–S6外加一个目标变量Y——基线后一年疾病进展的定量度量。文件为制表符分隔tab-delimited的 CSV 格式这一点直接影响两种实现里的解析配置。数据源获取CSV 的按需下载两种实现共用 utils.rs 中的download_csv_if_missing()pub fn download_csv_if_missing() - PathBuf { // Point file to current example directory let example_dir Path::new(env!(CARGO_MANIFEST_DIR)); let file_name example_dir.join(diabetes.csv); if file_name.exists() { println!(File already downloaded at {file_name:?}); } else { // Get file from web let url https://www4.stat.ncsu.edu/~boos/var.select/diabetes.tab.txt; let mut response reqwest::blocking::get(url).unwrap(); let mut file File::create(file_name).unwrap(); copy(mut response, mut file).unwrap(); }; file_name }几个值得注意的实现细节文件落盘位置由env!(CARGO_MANIFEST_DIR)决定即示例 crate 自身目录下的diabetes.csv因此下载发生在示例目录内不依赖运行时的工作目录使用reqwest的 blocking 客户端做一次性下载文件已存在时跳过保证可重复运行Cargo.toml 中reqwest启用了blockingfeatureCSV 解析依赖csvcrate 与带derive的serde。记录结构体用 serde 完成表头到字段的映射CSV 的表头是大写且语义不直观的AGE、SEX、S1…Y因此 diabetes_patient.rs 为每个字段手动指定serde(rename ...)并将字段收敛到能表达实际取值范围的最小类型#[derive(Deserialize, Serialize, Debug, Clone)] pub struct DiabetesPatient { /// Age in years #[serde(rename AGE)] pub age: i8, /// Sex categorical label #[serde(rename SEX)] pub sex: i8, /// Body mass index #[serde(rename BMI)] pub bmi: f32, /// Average blood pressure #[serde(rename BP)] pub bp: f32, /// S1: total serum cholesterol #[serde(rename S1)] pub tc: i16, // S2..S5 为 f32S6 (glu) 为 i8此处从略 /// Y: quantitative measure of disease progression one year after baseline #[serde(rename Y)] pub response: i16, }设计要点字段类型选择i8/i16/f32与取值范围匹配内存友好也体现了深度学习场景中通常以f32作为数值主类型的一致性S1–S6被重命名为可读字段tc、ldl、hdl、tch、ltg、gluY映射为response使下游训练代码不必再感知原始表头该结构体同时是后面两条实现路线的“交换格式”——无论走csv还是走 Polars最终都解码为DiabetesPatient。路线一InMemDatasetcsv serde 全量入内存dataset.rs 展示了最直接的实现方式pub struct DiabetesDataset { dataset: InMemDatasetDiabetesPatient, } impl DiabetesDataset { pub fn new() - ResultSelf, std::io::Error { // Download dataset csv file let path download_csv_if_missing(); // Build dataset from csv with tab (\t) delimiter let mut rdr csv::ReaderBuilder::new(); let rdr rdr.delimiter(b\t); let dataset InMemDataset::from_csv(path, rdr).unwrap(); let dataset Self { dataset }; Ok(dataset) } } // Implement the Dataset trait which requires get and len impl DatasetDiabetesPatient for DiabetesDataset { fn get(self, index: usize) - ResultDiabetesPatient, DatasetError { self.dataset.get(index) } fn len(self) - usize { self.dataset.len() } }from_csv 的底层实现InMemDataset::from_csv定义在 crates/burn-dataset/src/dataset/in_memory.rs 中其文档注释说明了两个关键约束你传入的csv::ReaderBuilder可以按自己的文件格式自由配置本例用它把分隔符改为b\t这正是 tab 分隔文件能正确解析的原因通过 serde 路径支持的字段类型为 String、整数、浮点数和布尔值。从源码结构看InMemDataset本体就是items: VecI见 in_memory.rs其Dataset实现的get对越界索引直接 paniclen返回items.len()。也就是说from_csv在构造期就把整个文件反序列化进一个Vec之后的随机访问是 O(1) 的内存读取代价是加载时的一次性内存占用。运行cargo run --example custom-csv-dataset入口 custom-csv-dataset.rs 做了三件事加载数据集并打印总行数应为 442随后分别取第 0 条与第 441 条首尾记录打印Debug输出用于直观验证解析结果let dataset DiabetesDataset::new().expect(Could not load diabetes dataset); println!(Dataset loaded with {} rows, dataset.len()); let item dataset.get(0).unwrap(); println!(First item:\n{item:?}); let item dataset.get(441).unwrap(); println!(Last item:\n{item:?});路线二DataframeDatasetPolars 列式加载dataframe_dataset.rs 演示了以 Polars 为后端的DataframeDatasetREADME 明确指出这种方式“更适合对更大规模数据集做高效的数据操作与分析”。两段式类型设计schema_type 与 cast_type核心是一张列定义表每个元素为(列名, 解析用类型, 最终输出类型)// Column definitions: (name, schema_type for parsing, cast_type for final output) const COLS: [(str, DataType, DataType)] [ (AGE, DataType::Int64, DataType::Int8), (SEX, DataType::Int64, DataType::Int8), (BMI, DataType::Float64, DataType::Float32), (BP, DataType::Float64, DataType::Float32), (S1, DataType::Int64, DataType::Int16), // S2..S5 为 (Float64, Float32)S6 为 (Int64, Int8)此处从略 (Y, DataType::Int64, DataType::Int16), ];从源码结构看这个“两段式”设计是有意为之CSV 解析阶段统一用 64 位宽类型Int64/Float64读取文本数值规避文本到窄类型的解析不确定性数据全部进入 DataFrame 后再逐列cast到与DiabetesPatient字段一致的窄类型既控制了内存占用又让 DataFrame 中的数值布局与 serde 结构体的期望完全对齐。惰性读取与列裁剪加载流程完整代码节选自 dataframe_dataset.rs// Build Schema let schema Schema::from_iter( COLS.iter() .map(|(name, schema_type, _)| Field::new((*name).into(), schema_type.clone())), ); let mut df LazyCsvReader::new(PlPath::new(path.to_str().unwrap())) .with_has_header(true) .with_separator(b\t) // 与路线一相同tab 分隔 .with_schema(Some(Arc::new(schema))) .finish()? .collect()?; // cast columns for (col, _, cast_type) in COLS { df.with_column(df.column(col)?.cast(cast_type)?.clone())?; } let dataset DataframeDataset::new(df)?;与路线一的逐行 serde 解码不同这里通过LazyCsvReader声明列名、有表头、制表符分隔以及显式 schema由 Polars 的引擎一次性完成列式扫描与类型转换对大文件而言 I/O 与内存行为都更可控。DataframeDataset 如何把行映射回结构体DataframeDataset定义在 crates/burn-dataset/src/dataset/dataframe.rspub struct DataframeDatasetI { df: DataFrame, len: usize, column_name_mapping: Vecusize, phantom: PhantomDataI, }构造时它会取df.height()作为len通过extract_field_names::I()基于I的 serde 结构信息得到字段名列表再把每个字段名映射到 DataFrame schema 中的列下标column_name_mapping。get(index)则用df.get_row(index)取出该行并反序列化为I。值得注意的是它复用了 serde 的字段名提取逻辑因此在 diabetes_patient.rs 中写的rename映射在这里同样生效——这就是两条路线能用同一个记录结构体的原因。DataframeDataset的get返回的是DataframeDatasetError而示例需要适配到 Burn 的DatasetError所以 trait 实现里做了一层错误转换dataframe_dataset.rsimpl DatasetDiabetesPatient for DiabetesDataframeDataset { fn get(self, index: usize) - ResultDiabetesPatient, DatasetError { self.dataset.get(index).map_err(DatasetError::new) } fn len(self) - usize { self.dataset.len() } }DatasetError::new接受任意std::error::Error Send Sync static的错误包装见 crates/burn-dataset/src/dataset/error.rs因此任何第三方数据集库的错误都可以这样平滑接入。运行需要启用 featurecargo run --example dataframe-dataset --features dataframe--features dataframe不可省略Cargo.toml 中定义了 feature 依赖关系且dataframe-dataset这个 example 被标记为required-features [dataframe][features] default [burn/dataset] dataframe [dep:burn-dataset, polars] # Dataframe support (optional) polars { workspace true, optional true, features [csv, temporal] } [[example]] name dataframe-dataset required-features [dataframe]也就是说默认 feature 只开burn/dataset覆盖csv serde 路线打开dataframe后会额外引入burn-dataset启用其dataframefeature和带csv、temporalfeature 的polarssrc/lib.rs 里dataframe_dataset模块同样受#[cfg(feature dataframe)]门控不开 feature 时该模块不参与编译。Dataset trait两条路线的共同契约Burn 的Datasettrait 是最小契约只需实现get(index)与len()。从 crates/burn-dataset/src/dataset/base.rs 看trait 还带有基于逐项get的get_many默认实现并对越界索引断言 panic且对Arcdyn Dataset、BoxD等包装类型有透明的转发实现这意味着自定义数据集可以在不改变接口的情况下被智能指针包装、共享或动态分发。InMemDataset还提供了from_dataset把任意Dataset迭代一遍收进内存与from_json_rows逐行 JSONL 反序列化等构造器见 in_memory.rs说明“先有任意 Dataset 实现、再一键转为内存版”也是仓库支持的组合方式。两条路线的选型建议结合示例代码与源码可以归纳出清晰的适用边界维度InMemDatasetcsv serdeDataframeDatasetPolars实现复杂度低一个ReaderBuilderfrom_csv即可中需要定义 schema 与 cast 流程内存行为构造期一次性全量反序列化进Vec列式存储解析与窄化分离对宽/大表更友好依赖面csvserde默认 feature 即含需--features dataframe引入polars适用场景中小规模、结构简单、快速起步大规模数据、需要列式操作与分析的管线两者的共同点是最终都收敛到同一个DiabetesPatient记录结构体、都实现了同一个Datasettrait——因此对训练代码而言二者完全可互换区别只在于数据加载层的工程取舍。小结examples/custom-csv-dataset用一份 442 行的 diabetes 数据完整演示了 Burn 自定义 CSV 数据集的两个标准动作记录建模用带serde(rename)的结构体把表头映射为类型收敛的 Rust 字段src/diabetes_patient.rs加载实现小数据走InMemDataset::from_csv记得用ReaderBuilder适配制表符分隔大数据走 PolarsLazyCsvReaderDataframeDatasetsrc/dataframe_dataset.rs契约接入只需实现get/len即满足Datasettrait即可无缝进入 Burn 的训练循环。运行入口为cargo run --example custom-csv-dataset与cargo run --example dataframe-dataset --features dataframe两个示例都会打印首尾两条记录用于肉眼校验解析正确性。【免费下载链接】burnBurn is a next generation tensor library and Deep Learning Framework that doesnt compromise on flexibility, efficiency and portability.项目地址: https://gitcode.com/GitHub_Trending/bu/burn创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

关于本文作者

来自尧图内容编辑团队

尧图内容编辑团队 内容团队

尧图内容编辑团队

本文由尧图网络内容编辑团队执笔。团队由资深项目经理、前端工程师与设计师组成,所有内容均来自亲手交付的真实项目,先讲清问题、再给出可落地的解法。尧图深耕北京网站建设十年,服务过京华建材集团、智造科技等各行业客户,把一线经验沉淀为可复用的行业观察。

  • 十年建站经验,覆盖建材、制造、服务、文创等
  • 项目经理把关选题与事实准确性
  • 工程师与设计师联合撰写专业细节
  • 统一编辑规范,保证文风与排版一致
  • 每月复盘转化数据,迭代选题方向

延伸阅读

相关资讯与近期热门内容

深度阅读推荐

建站决策前值得细读的三篇

网站改版的5个关键决策
2024-08-12

网站改版的5个关键决策

什么时候该改版、改到什么程度、如何避免流量掉光,京华建材集团改版复盘给出答案。

获取专属建站方案

看完文章,把您的行业与预算告诉我们,免费获取一份量身定制的官网建设方案与报价。

立即免费咨询