mirror of
https://github.com/p-e-w/heretic.git
synced 2026-10-02 08:51:27 -07:00
feat: A better logic for detecting the thinking prefix in the response. (#423)
* fix: Check response prefix. In case if the model adds <think> at the end of user prompt and only generates </think> end, then the detection goes wrong. * fix: Handle whitespace in CoT, remove redundant checks * fix: Type checker * fix: try looking at tag positions, not text end regex * ruff, fix import sorting * fix: rich markup * fix: I missed other ones * feat: Update SHA256SUMS file hashes in the tests. A major change that affects reproducibility. * fix: It is now sensible to also update the extra SHA256SUMS.ci2 file. * fix: Only consider whitespace, no other text or instructions. because mistral-3 as additional reasoning instructions in its chat template. And I suppose many other models can have it too. * fix: Update windows hashes. * fix: Update CI hashes too * as always update case two of mistral-3 (ci2 hash) * docs: update comment * feat: Handle the edge case for models having additional instructions. * fix: Update windows hash for mistral-3 * docs: remove a line from the comments because I'm not sure about GPT-OSS models' thinking tags and it cannot be confirmed using an untrained tiny GPT-OSS model. And inference fallback would of course generate gibberish as the model cannot understand additional instructions about 'how to generate response and how to think' from the chat_template. * fix: a few things. * fix: Update hash for qwen3.5 after the whitespace fix for its response prefix. * docs: Update comment * fix: Update qwen3.5 hash for CI * fix: Remove Case 2 which only serves tests unnecessary * fix: Hash * fix: concern is valid enough, so we use a small text. add a comment too * docs: minor
This commit is contained in:
+78
-28
@@ -39,13 +39,14 @@ import logging
|
||||
import math
|
||||
import os
|
||||
import random
|
||||
import re
|
||||
import time
|
||||
import warnings
|
||||
from dataclasses import asdict
|
||||
from importlib.metadata import version
|
||||
from os.path import commonprefix
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from typing import Any, cast
|
||||
|
||||
import huggingface_hub
|
||||
import lm_eval
|
||||
@@ -65,6 +66,7 @@ from optuna.storages.journal import JournalFileBackend, JournalFileOpenLock
|
||||
from optuna.trial import FrozenTrial, TrialState, create_trial
|
||||
from pydantic import ValidationError
|
||||
from questionary import Choice, Style
|
||||
from rich.markup import escape
|
||||
from rich.table import Table
|
||||
from rich.text import Text
|
||||
from rich.traceback import install
|
||||
@@ -477,40 +479,88 @@ def run():
|
||||
print()
|
||||
print("Checking for common response prefix...")
|
||||
prefix_check_prompts = good_prompts[:100] + bad_prompts[:100]
|
||||
responses = model.get_responses_batched(prefix_check_prompts)
|
||||
|
||||
# Despite being located in os.path, commonprefix actually performs
|
||||
# a naive string operation without any path-specific logic,
|
||||
# which is exactly what we need here. Trailing spaces are removed
|
||||
# to avoid issues where multiple different tokens that all start
|
||||
# with a space character lead to the common prefix ending with
|
||||
# a space, which would result in an uncommon tokenization.
|
||||
settings.response_prefix = commonprefix(responses).rstrip(" ")
|
||||
# Detect if the model's chat template inserts a reasoning tag on its own
|
||||
# at the end of user's prompt (e.g. <think>) by using a dummy prompt.
|
||||
# If found, then we use the full closed CoT as the response prefix.
|
||||
# LiquidAI's LFM models do this (Lfm2ForCausalLM).
|
||||
|
||||
if settings.response_prefix:
|
||||
print(f"* Prefix found: [bold]{settings.response_prefix!r}[/]")
|
||||
# This cast is valid because str is the return type
|
||||
# for a single chat operation with tokenize=False.
|
||||
dummy_prompt = cast(
|
||||
str,
|
||||
model.tokenizer.apply_chat_template(
|
||||
[{"role": "user", "content": "This is a dummy prompt."}],
|
||||
add_generation_prompt=True,
|
||||
tokenize=False,
|
||||
),
|
||||
)
|
||||
|
||||
for cot_initializer, closed_cot_block in settings.chain_of_thought_skips:
|
||||
if settings.response_prefix.startswith(cot_initializer):
|
||||
settings.response_prefix = closed_cot_block
|
||||
print(
|
||||
f"* Closed Chain-of-Thought block: [bold]{settings.response_prefix!r}[/]"
|
||||
)
|
||||
cot_skip_applied = False
|
||||
|
||||
# When using a Chain-of-Thought skip, we need to check that the prefix
|
||||
# is actually complete (e.g. not missing a trailing newline).
|
||||
print("* Rechecking with prefix...")
|
||||
responses = model.get_responses_batched(prefix_check_prompts)
|
||||
additional_prefix = commonprefix(responses).rstrip(" ")
|
||||
if additional_prefix:
|
||||
settings.response_prefix += additional_prefix
|
||||
for cot_initializer, closed_cot_block in settings.chain_of_thought_skips:
|
||||
# Match the tag and ignore any whitespace characters following it at the end
|
||||
# (if any), including spaces, tabs, and linebreaks. This is required for models
|
||||
# having whitespaces after the tags.
|
||||
pattern = rf"{re.escape(cot_initializer)}\s*$"
|
||||
match = re.search(pattern, dummy_prompt)
|
||||
|
||||
if match:
|
||||
# We use only the closed CoT block here. Any whitespaces
|
||||
# will be handled by the 'Rechecking with prefix' logic below.
|
||||
settings.response_prefix = closed_cot_block
|
||||
print(
|
||||
f"* Closed Chain-of-Thought block: [bold]{escape(repr(settings.response_prefix))}[/]"
|
||||
)
|
||||
cot_skip_applied = True
|
||||
break
|
||||
|
||||
# Fallback to inference for models like mistral-3 which are specifically
|
||||
# instructed to generate thinking tags using the system prompt in their
|
||||
# chat template, instead of inserting a prefix tag (e.g. <think>) at
|
||||
# the end of user prompt like the case above. We expect the model to
|
||||
# generate those tags.
|
||||
if settings.response_prefix is None:
|
||||
responses = model.get_responses_batched(prefix_check_prompts)
|
||||
|
||||
# Despite being located in os.path, commonprefix actually performs
|
||||
# a naive string operation without any path-specific logic,
|
||||
# which is exactly what we need here. Trailing spaces are removed
|
||||
# to avoid issues where multiple different tokens that all start
|
||||
# with a space character lead to the common prefix ending with
|
||||
# a space, which would result in an uncommon tokenization.
|
||||
settings.response_prefix = commonprefix(responses).rstrip(" ")
|
||||
|
||||
if settings.response_prefix:
|
||||
print(
|
||||
f"* Prefix found: [bold]{escape(repr(settings.response_prefix))}[/]"
|
||||
)
|
||||
|
||||
for (
|
||||
cot_initializer,
|
||||
closed_cot_block,
|
||||
) in settings.chain_of_thought_skips:
|
||||
if settings.response_prefix.startswith(cot_initializer):
|
||||
settings.response_prefix = closed_cot_block
|
||||
print(
|
||||
f"* Extended prefix found: [bold]{settings.response_prefix!r}[/]"
|
||||
f"* Closed Chain-of-Thought block: [bold]{escape(repr(settings.response_prefix))}[/]"
|
||||
)
|
||||
cot_skip_applied = True
|
||||
break
|
||||
else:
|
||||
print("* None found")
|
||||
|
||||
break
|
||||
else:
|
||||
print("* None found")
|
||||
if cot_skip_applied:
|
||||
# When using a Chain-of-Thought skip, we need to check that the prefix
|
||||
# is actually complete (e.g. not missing a trailing newline).
|
||||
print("* Rechecking with prefix...")
|
||||
responses = model.get_responses_batched(prefix_check_prompts)
|
||||
additional_prefix = commonprefix(responses).rstrip(" ")
|
||||
if additional_prefix:
|
||||
settings.response_prefix += additional_prefix
|
||||
print(
|
||||
f"* Extended prefix found: [bold]{escape(repr(settings.response_prefix))}[/]"
|
||||
)
|
||||
|
||||
evaluator = Evaluator(settings, model)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user