diff --git a/scripts/prepare_hidden_states.py b/scripts/prepare_hidden_states.py index 109c014..89d254f 100644 --- a/scripts/prepare_hidden_states.py +++ b/scripts/prepare_hidden_states.py @@ -814,16 +814,20 @@ def main(): with rank_0_priority(): print_with_rank("Loading/building dataset cache...") - dataset = Dataset.from_generator( - generator=safe_conversations_generator, - gen_kwargs={"file_path": args.data_path}, - cache_dir=os.path.join( - os.path.dirname(os.path.dirname(os.path.abspath(__file__))), - "cache", - "hf_dataset", - ), - num_proc=min(args.build_dataset_num_proc, 32), - ) + if args.is_preformatted: + # Preserve the text column: conversation normalization discards it. + dataset = Dataset.from_json(args.data_path, cache_dir=args.cache_dir) + else: + dataset = Dataset.from_generator( + generator=safe_conversations_generator, + gen_kwargs={"file_path": args.data_path}, + cache_dir=os.path.join( + os.path.dirname(os.path.dirname(os.path.abspath(__file__))), + "cache", + "hf_dataset", + ), + num_proc=min(args.build_dataset_num_proc, 32), + ) if args.num_samples is not None and capture_plan.loss_mask_filter is None: dataset = dataset.select(range(args.num_samples)) # Tokenizer and cache key