@@ -297,8 +297,16 @@ def save_dataset(dataset: DashAIDataset, path: Union[str, os.PathLike]) -> None:
297297 writer .close ()
298298
299299 metadata_filepath = os .path .join (path , "splits.json" )
300- with open (metadata_filepath , "w" ) as f :
301- json .dump (dataset .splits , f , indent = 2 , sort_keys = True , ensure_ascii = False )
300+ # Update splits with dataset shape and column names
301+ metadata = dataset .splits
302+ metadata .update (
303+ {
304+ "total_rows" : dataset .shape [0 ],
305+ "column_names" : dataset .column_names ,
306+ }
307+ )
308+ with open (metadata_filepath , "w" , encoding = "utf-8" ) as f :
309+ json .dump (metadata , f , indent = 2 , sort_keys = True , ensure_ascii = False )
302310
303311
304312@beartype
@@ -767,24 +775,16 @@ def get_dataset_info(dataset_path: str) -> object:
767775 else :
768776 splits_data = {"split_indices" : {}}
769777
770- data_filepath = os .path .join (dataset_path , "data.arrow" )
771- with pa .OSFile (data_filepath , "rb" ) as source :
772- reader = ipc .open_file (source )
773- schema = reader .schema
774- column_names = schema .names
775-
776- total_rows = 0
777- for i in range (reader .num_record_batches ):
778- total_rows += reader .get_batch (i ).num_rows
779-
780778 splits = splits_data .get ("split_indices" , {})
781779 train_indices = splits .get ("train" , [])
782780 test_indices = splits .get ("test" , [])
783781 val_indices = splits .get ("validation" , [])
782+ total_rows = splits_data .get ("total_rows" , 0 )
783+ column_names = splits_data .get ("column_names" , [])
784784
785785 return {
786786 "total_rows" : total_rows ,
787- "total_columns" : len (schema ),
787+ "total_columns" : len (column_names ),
788788 "column_names" : column_names ,
789789 "train_size" : len (train_indices ),
790790 "test_size" : len (test_indices ),
0 commit comments