From ff956aedbf61c101a590e035d5faf853c9d36f90 Mon Sep 17 00:00:00 2001 From: nandanadileep Date: Fri, 12 Jun 2026 09:08:10 +0000 Subject: [PATCH] Fix #2407: Fix get_wikitext2 tokenization bug causing sequence length w Signed-off-by: nandanadileep --- optimum/gptq/data.py | 10 ++++++---- 1 file changed, 6 insertions(+), 4 deletions(-) diff --git a/optimum/gptq/data.py b/optimum/gptq/data.py index 127e6676cd..baf6ee2397 100644 --- a/optimum/gptq/data.py +++ b/optimum/gptq/data.py @@ -125,17 +125,19 @@ def get_wikitext2(tokenizer: Any, seqlen: int, nsamples: int, split: str = "trai data = load_dataset("Salesforce/wikitext", "wikitext-2-raw-v1", split="train") elif split == "validation": data = load_dataset("Salesforce/wikitext", "wikitext-2-raw-v1", split="test") - # length of 288059 should be enough - text = "".join([" \n" if s == "" else s for s in data["text"][:1000]]) - - enc = tokenizer(text, return_tensors="pt") dataset = [] for _ in range(nsamples): + while True: + i = random.randint(0, len(data) - 1) + enc = tokenizer(data[i]["text"], return_tensors="pt") + if enc.input_ids.shape[1] >= seqlen: + break i = random.randint(0, enc.input_ids.shape[1] - seqlen - 1) j = i + seqlen inp = enc.input_ids[:, i:j] attention_mask = torch.ones_like(inp) dataset.append({"input_ids": inp, "attention_mask": attention_mask}) + return dataset