diff --git a/unsloth-cli.py b/unsloth-cli.py index 5135d00493..ef8167ad39 100644 --- a/unsloth-cli.py +++ b/unsloth-cli.py @@ -103,19 +103,13 @@ def run(args): from transformers.utils import strtobool if args.raw_text_file: - # Use raw text loader - returns pre-tokenized data + # Use raw text loader loader = RawTextDataLoader(tokenizer, args.chunk_size, args.stride) - dataset = loader.load_from_file(args.raw_text_file, return_tensors = True) - # Mark dataset as pre-tokenized to skip text formatting - dataset._is_pretokenized = True - return dataset + dataset = loader.load_from_file(args.raw_text_file) elif args.dataset.endswith((".txt", ".md", ".json", ".jsonl")): - # Auto-detect local raw text files - returns pre-tokenized data - loader = RawTextDataLoader(tokenizer, args.chunk_size, args.stride) - dataset = loader.load_from_file(args.dataset, return_tensors = True) - # Mark dataset as pre-tokenized to skip text formatting - dataset._is_pretokenized = True - return dataset + # Auto-detect local raw text files + loader = RawTextDataLoader(tokenizer) + dataset = loader.load_from_file(args.dataset) else: # Check for modelscope usage use_modelscope = strtobool( @@ -129,9 +123,9 @@ def run(args): # Existing HuggingFace dataset logic dataset = load_dataset(args.dataset, split = "train") - # Apply formatting for structured datasets (text-based) + # Apply formatting for structured datasets dataset = dataset.map(formatting_prompts_func, batched = True) - return dataset + return dataset # Load dataset using smart loader dataset = load_dataset_smart(args) @@ -427,19 +421,5 @@ if __name__ == "__main__": "--stride", type = int, default = 512, help = "Overlap between chunks" ) - TRAINING_MODES = { - "instruction": "Standard instruction-following", - "causal": "Causal language modeling (raw text)", - "completion": "Text completion tasks", - } - - parser.add_argument( - "--training_mode", - type = str, - default = "instruction", - choices = list(TRAINING_MODES.keys()), - help = "Training mode for the model", - ) - args = parser.parse_args() run(args)