Flatten _text_run_ids output (input_ids is [1,seq])

This commit is contained in:
Daniel Maddern 2026-08-19 22:54:32 +07:00
parent 428c83d1d8
commit 0903b849b2

View file

@ -384,7 +384,9 @@ def _text_run_ids(prompt: str) -> list[int]:
tokenizer_dir = Path(__file__).with_name("qwen25_tokenizer")
if not tokenizer_dir.exists():
raise FileNotFoundError(f"Qwen tokenizer directory missing: {tokenizer_dir}")
return H3PromptTokenizer(tokenizer_dir)(prompt or " ")
ids = H3PromptTokenizer(tokenizer_dir)(prompt or " ")
# input_ids is [1, seq]; flatten to a Python list of ints.
return [int(t) for t in ids.reshape(-1).tolist()]
@dataclass