diff --git a/mlx_lm/quant/dwq.py b/mlx_lm/quant/dwq.py index 769327d..a0796af 100644 --- a/mlx_lm/quant/dwq.py +++ b/mlx_lm/quant/dwq.py @@ -42,13 +42,11 @@ def compute_dwq_targets( if rank == 0: path = path / split path.mkdir(parents=True, exist_ok=True) - for i, (batch, _) in ( - pbar := tqdm( - enumerate(iterate_batches(data, batch_size, max_seq_length, seed=seed)), - total=len(data) // batch_size, - desc=f"Computing targets for {split}", - disable=rank != 0, - ) + for i, (batch, _) in tqdm( + enumerate(iterate_batches(data, batch_size, max_seq_length, seed=seed)), + total=len(data) // batch_size, + desc=f"Computing targets for {split}", + disable=rank != 0, ): batch = batch[:, :-1] logits = model(batch) @@ -216,15 +214,24 @@ def load_data( max_seq_length: int, num_valid_samples: int = 32, ): - args = types.SimpleNamespace( - hf_dataset={ - "path": data_path, - "train_split": "train", - "valid_split": "train[:1]", - }, - train=True, - test=False, - ) + if Path(data_path).exists(): + args = types.SimpleNamespace( + data=data_path, + hf_dataset=False, + train=True, + test=False, + mask_prompt=False, + ) + else: + args = types.SimpleNamespace( + hf_dataset={ + "path": data_path, + "train_split": "train", + "valid_split": "train[:1]", + }, + train=True, + test=False, + ) dataset = load_dataset(args, tokenizer)[0] perm = np.random.permutation(len(dataset)) train_perm = perm[:num_samples].tolist() @@ -328,7 +335,10 @@ def main(): has_targets = False target_dir = None - tokenizer = load_tokenizer(args.model) + tokenizer = load_tokenizer( + args.model, + tokenizer_config_extra={"trust_remote_code": args.trust_remote_code}, + ) train_data, valid_data = load_data( tokenizer, args.data_path, args.num_samples, args.max_seq_length