diff --git a/examples/speculative_decoding/collect_hidden_states/compute_hidden_states_hf.py b/examples/speculative_decoding/collect_hidden_states/compute_hidden_states_hf.py index 1060df45622..713b2326fcc 100644 --- a/examples/speculative_decoding/collect_hidden_states/compute_hidden_states_hf.py +++ b/examples/speculative_decoding/collect_hidden_states/compute_hidden_states_hf.py @@ -206,10 +206,12 @@ async def submit_generates(): continue # Tokenize and check length - tokenized = tokenizer.apply_chat_template( - conversations, return_tensors="pt", add_generation_template=False - ) - input_ids = tokenized["input_ids"] if isinstance(tokenized, dict) else tokenized + # return_dict=True ensures BatchEncoding is returned on all transformers + # versions: in <5.0 the default is False (returns raw tensor), in 5.0+ + # the default changed to True (returns BatchEncoding). + input_ids = tokenizer.apply_chat_template( + conversations, return_tensors="pt", return_dict=True, add_generation_template=False + )["input_ids"] num_input_tokens = input_ids.shape[1] if num_input_tokens <= 10 or num_input_tokens > args.max_seq_len: num_skipped_too_long += 1