diff --git a/README.md b/README.md index 9bc17b2..b539964 100644 --- a/README.md +++ b/README.md @@ -16,41 +16,95 @@ import torch from transformers import AutoModelForCausalLM, AutoTokenizer from repeng import ControlVector, ControlModel, DatasetEntry +from repeng.utils import make_dataset + +# load and wrap model +model_name = "mistralai/Mistral-7B-Instruct-v0.3" + +# If you need quantization, but can lead to issues. For +# example Gemma3 models seem to silently generate empty strings. +# from transformers import BitsAndBytesConfig +# bnb_config = BitsAndBytesConfig( +# load_in_4bit=True, +# bnb_4bit_quant_type="nf4", +# bnb_4bit_compute_dtype=torch.bfloat16, +# bnb_4bit_use_double_quant=True, +# ) + +model = AutoModelForCausalLM.from_pretrained( + model_name, + # quantization_config=bnb_config, + # torch_dtype=torch.float16, + ) +) -# load and wrap Mistral-7B -model_name = "mistralai/Mistral-7B-Instruct-v0.1" -model = AutoModelForCausalLM.from_pretrained(model_name, torch_dtype=torch.float16) -model = ControlModel(model, list(range(-5, -18, -1))) +# wrap the model to give us control +model = ControlModel( + model, + # layer_ids=list(range(-5, -18, -1)) # specify layers to control by layer ID + layer_zones=[[0.3, 0.5]], # control layers with relative depth in [0.3, 0.5[ +) -def make_dataset(template: str, pos_personas: list[str], neg_personas: list[str], suffixes: list[str]): - # see notebooks/experiments.ipynb for a definition of `make_dataset` - ... +tokenizer = AutoTokenizer.from_pretrained( + model_name, + # quantization_config=bnb_config, +) # generate a dataset with closely-opposite paired statements trippy_dataset = make_dataset( - "Act as if you're extremely {persona}.", - ["high on psychedelic drugs"], - ["sober from psychedelic drugs"], - truncated_output_suffixes, + # you can use either chat as dicts... + template=[ + {"role": "system", "content": "You talk like you are {persona}."}, + {"role": "user", "content": "{suffix}"}, + ], + # ...or directly strings: + # template="Act as if you're {persona}. Someone comes at you and says '{suffix}'.", + + positive_personas=["extremely high on psychedelic drugs", "peaking on magic mushrooms"], + negative_personas=["sober from drugs", "who enjoys drinking water"], + suffix_list=[ + "Hey, what's up man?", + "Hey, what's up girl?", + "Welcome Mr Musk, come this way.", + "How have you been feeling lately with the medications?", + ], ) # train the vector—takes less than a minute! trippy_vector = ControlVector.train(model, tokenizer, trippy_dataset) +# Now we must give the scenario for the generation we will engineer: +scenario: str = tokenizer.apply_chat_template( + conversation=[ + { + "role": "user", + "content": "Give me a one-sentence pitch for a TV show." + }, + ], + continue_final_message=False, + tokenize=False, +) + +# Or directly as a str +# scenario=f"[INST] Give me a one-sentence pitch for a TV show. [/INST]", + # set the control strength and let inference rip! for strength in (-2.2, 1, 2.2): print(f"strength={strength}") model.set_control(trippy_vector, strength) out = model.generate( **tokenizer( - f"[INST] Give me a one-sentence pitch for a TV show. [/INST]", + scenario, return_tensors="pt" - ), + ).to(model.device), do_sample=False, - max_new_tokens=128, + # temperature=1.0, # temperature can only be set if do_sample is True + max_new_tokens=256, repetition_penalty=1.1, ) print(tokenizer.decode(out.squeeze()).strip()) + # or if you want to display the special tokens: + # print(tokenizer.decode(out.squeeze(), skip_special_tokens=False).strip()) print() ``` @@ -69,6 +123,7 @@ For a more detailed explanation of how the library works and what it can do, see * For a list of changes by version, see the [CHANGELOG](https://github.com/vgel/repeng/blob/main/CHANGELOG). * For quantized use, you may be interested in [llama.cpp#5970](https://github.com/ggerganov/llama.cpp/pull/5970)—after training a vector with `repeng`, export it by calling `vector.export_gguf(filename)` and then use it in `llama.cpp` with any quant! +* To load gguf files directly, you can run into OOM errors, see [this github issue for more](See here: https://github.com/huggingface/transformers/issues/34417). * Vector training *currently does not work* with MoE models (such as Mixtral). (This is theoretically fixable with some work, let me know if you're interested.) * Some example notebooks require `accelerate`, which must be manually installed with `pip install accelerate`. (This can also be done in the notebook with the IPython magic `%pip install accelerate`.) diff --git a/pyproject.toml b/pyproject.toml index ab542e9..d08aea9 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -12,8 +12,11 @@ dependencies = [ "transformers>=4.36.2", "tqdm>=4.66.1", "gguf>=0.13.0", + "loguru>=0.7.3", + ] + [dependency-groups] dev = ["pytest>=8.0.2", "ruff>=0.8.3"] diff --git a/repeng/control.py b/repeng/control.py index 6cf6786..98d4349 100644 --- a/repeng/control.py +++ b/repeng/control.py @@ -5,6 +5,8 @@ import torch from transformers import PretrainedConfig, PreTrainedModel +from repeng.utils import get_num_hidden_layer + if typing.TYPE_CHECKING: from .extract import ControlVector @@ -16,20 +18,68 @@ class ControlModel(torch.nn.Module): A wrapped language model that can have controls set on its layers with `self.set_control`. """ - def __init__(self, model: PreTrainedModel, layer_ids: typing.Iterable[int]): + def __init__( + self, + model: PreTrainedModel, + layer_ids: typing.Optional[typing.Iterable[int]]=None, + layer_zones: typing.Optional[typing.Iterable[float]]=None, + ): """ **This mutates the wrapped `model`! Be careful using `model` after passing it to this class.** Build a new ControlModel around a model instance, initializing control on - the layers specified in `layer_ids`. + the layers specified in `layer_ids` or `layer_zones`. + + To control layers #3 and 5, use layer_ids=[3,5]. + To control layers by their relative depth, use layer_zones=[[0.1, 0.5]] to + control the layers with depth between 10% and 50% (left inclusive). you + can specify multiple zones but no overlapping nor empty zones are allowed. """ + assert (layer_ids or layer_zones) and not (layer_ids and layer_zones), "Must supply either layer_ids or layer_zones argument" + super().__init__() self.model = model + # Get the number of layers + layer_ids = list(range(get_num_hidden_layer(model))) + nlayers = len(layer_ids) + + if layer_zones: + self.layer_ids = [] + for start_zone, end_zone in layer_zones: + assert ( + start_zone < end_zone + and start_zone >= 0 + and start_zone <= 1 + and end_zone >= 0 + and end_zone <= 1 + ), "wrong layer_zones format" + if end_zone != 1.0: + new_layers = [ + ilayer + for ilayer in layer_ids + if start_zone <= (ilayer / nlayers) < end_zone + ] + else: # trick to make sure to include the last layers if desired + new_layers = [ + ilayer + for ilayer in layer_ids + if start_zone <= (ilayer / nlayers) + ] + assert new_layers, f"No layers found in zone {start_zone} to {end_zone}" + assert not any(nl in self.layer_ids for nl in new_layers), "Overlapping zones found" + self.layer_ids.extend(new_layers) + else: + # remap to make sure they are not negative + self.layer_ids = layer_ids + + assert self.layer_ids, "No layers to control" + layers = model_layer_list(model) - self.layer_ids = [i if i >= 0 else len(layers) + i for i in layer_ids] - for layer_id in layer_ids: + assert len(layers) == len(layer_ids) + + for layer_id in self.layer_ids: layer = layers[layer_id] if not isinstance(layer, ControlModule): layers[layer_id] = ControlModule(layer) @@ -164,7 +214,8 @@ def forward(self, *args, **kwargs): assert len(control.shape) == len(modified.shape) control = control.to(modified.device) - norm_pre = torch.norm(modified, dim=-1, keepdim=True) + if self.params.normalize: + norm_pre = torch.norm(modified, dim=-1, keepdim=True) # we should ignore the padding tokens when doing the activation addition # mask has ones for non padding tokens and zeros at padding tokens. @@ -180,10 +231,10 @@ def forward(self, *args, **kwargs): .reshape(target_shape[0], target_shape[1], 1) ) mask = mask.to(modified.dtype).to(modified.device) + modified = self.params.operator(modified, control * mask) else: - mask = 1.0 + modified = self.params.operator(modified, control) - modified = self.params.operator(modified, control * mask) if self.params.normalize: norm_post = torch.norm(modified, dim=-1, keepdim=True) @@ -201,9 +252,15 @@ def model_layer_list(model: ControlModel | PreTrainedModel) -> torch.nn.ModuleLi if isinstance(model, ControlModel): model = model.model - if hasattr(model, "model"): # mistral-like - return model.model.layers + if hasattr(model, "language_model"): # gemmma3 like + layers = model.language_model.layers + elif hasattr(model, "layers"): # qwen3-like + layers = model.layers + elif hasattr(model, "base_model"): # mamba like + layers = model.base_model.layers elif hasattr(model, "transformer"): # gpt-2-like - return model.transformer.h + layers = model.transformer.h + elif hasattr(model, "model"): # mistral-like + layers = model.model.layers else: raise ValueError(f"don't know how to get layer list for {type(model)}") diff --git a/repeng/extract.py b/repeng/extract.py index 8d78a75..485ccf1 100644 --- a/repeng/extract.py +++ b/repeng/extract.py @@ -12,13 +12,9 @@ from .control import ControlModel, model_layer_list from .saes import Sae +from .utils import DatasetEntry, get_model_name, autocorrect_chat_templates, get_num_hidden_layer -@dataclasses.dataclass -class DatasetEntry: - positive: str - negative: str - @dataclasses.dataclass class ControlVector: @@ -255,14 +251,16 @@ def read_representations( Extract the representations based on the contrast dataset. """ if not hidden_layers: - hidden_layers = range(-1, -model.config.num_hidden_layers, -1) + hidden_layers = list(range(get_num_hidden_layer(model))) - # normalize the layer indexes if they're negative n_layers = len(model_layer_list(model)) - hidden_layers = [i if i >= 0 else n_layers + i for i in hidden_layers] # the order is [positive, negative, positive, negative, ...] - train_strs = [s for ex in inputs for s in (ex.positive, ex.negative)] + train_strs = autocorrect_chat_templates( + messages=[s for ex in inputs for s in (ex.positive, ex.negative)], + tokenizer=tokenizer, + model=model, + ) layer_hiddens = batched_get_hiddens( model, tokenizer, train_strs, hidden_layers, batch_size @@ -273,7 +271,7 @@ def read_representations( # get directions for each layer using PCA directions: dict[int, np.ndarray] = {} - for layer in tqdm.tqdm(hidden_layers): + for layer in tqdm.tqdm(hidden_layers, desc="Altering directions"): h = layer_hiddens[layer] assert h.shape[0] == len(inputs) * 2 @@ -343,10 +341,9 @@ def batched_get_hiddens( ] hidden_states = {layer: [] for layer in hidden_layers} with torch.no_grad(): - for batch in tqdm.tqdm(batched_inputs): + for batch in tqdm.tqdm(batched_inputs, desc="Computing activations"): # get the last token, handling right padding if present - encoded_batch = tokenizer(batch, padding=True, return_tensors="pt") - encoded_batch = encoded_batch.to(model.device) + encoded_batch = tokenizer(batch, padding=True, return_tensors="pt").to(model.device) out = model(**encoded_batch, output_hidden_states=True) attention_mask = encoded_batch["attention_mask"] for i in range(len(batch)): diff --git a/repeng/saes.py b/repeng/saes.py index bf394dc..4f2d9b7 100644 --- a/repeng/saes.py +++ b/repeng/saes.py @@ -71,7 +71,7 @@ def decode(self, features: np.ndarray) -> np.ndarray: huggingface_hub.snapshot_download(repo_id, revision=revision) ) layer_dict: dict[int, SaeLayer] = {} - for layer in tqdm.tqdm(layers): + for layer in tqdm.tqdm(layers, desc="Creating SAE"): eleuther_layer = layer - 1 # see docstr # this is in `sae` but to load the dtype we want, need to reimpl some stuff layer_path = base_path / f"layers.{eleuther_layer}" diff --git a/repeng/tests.py b/repeng/tests.py index d32cd77..4d60857 100644 --- a/repeng/tests.py +++ b/repeng/tests.py @@ -7,6 +7,7 @@ from . import ControlModel, ControlVector, DatasetEntry from .control import model_layer_list +from .utils import make_dataset def test_layer_list(): @@ -15,6 +16,35 @@ def test_layer_list(): _, lts = load_llama_tinystories_model() assert len(model_layer_list(lts)) == 4 +def test_make_dataset(): + trippy_dataset = make_dataset( + template=[ + {"role": "system", "content": "You talk like you are {persona}."}, + {"role": "user", "content": "{suffix}"}, + ], + # template="Act as if you're {persona}. Someone comes at you and says '{suffix}'.", + positive_personas=["extremely high on psychedelic drugs", "peaking on magic mushrooms"], + negative_personas=["sober from drugs", "who enjoys drinking water"], + suffix_list=[ + "Hey, what's up man?", + "Hey, what's up girl?", + "Welcome Mr Musk, come this way.", + "How have you been feeling lately with the medications?", + ], + ) + trippy_dataset = make_dataset( + template="Act as if you're {persona}. Someone comes at you and says '{suffix}'.", + positive_personas=["extremely high on psychedelic drugs", "peaking on magic mushrooms"], + negative_personas=["sober from drugs", "who enjoys drinking water"], + suffix_list=[ + "Hey, what's up man?", + "Hey, what's up girl?", + "Welcome Mr Musk, come this way.", + "How have you been feeling lately with the medications?", + ], + ) + + def test_round_trip_gguf(): tokenizer, model = load_llama_tinystories_model() @@ -74,7 +104,7 @@ def gen(vector: ControlVector | None, strength_coeff: float | None = None): assert happy == gen(happy_vector * 20) assert happy == gen(-(happy_vector * -20)) - assert sad == 'You are feeling the fucking damn goddamn worst,"' + assert sad in ['You are feeling the fucking damn goddamn worst,"', 'You are feeling the fucking damn goddamn fuck,"'] # these should be identical assert sad == gen(happy_vector, -50.0) assert sad == gen(happy_vector * -50) @@ -115,8 +145,8 @@ def gen(vector: ControlVector | None, strength_coeff: float | None = None): print(" cat:", cat) assert baseline.removeprefix(prompt) == " big, red" - assert mushroom.removeprefix(prompt) == " small plant." - assert cat.removeprefix(prompt) == " cat Bud guitar" + assert mushroom.removeprefix(prompt) in [" small plant.", " big cherry"] + assert cat.removeprefix(prompt) in [" cat Bud guitar", " guitar Bud guitar"] ################################################################################ @@ -169,25 +199,6 @@ def model_generate( return tokenizer.decode(out.squeeze()) # type: ignore -def make_dataset( - template: str, - positive_personas: list[str], - negative_personas: list[str], - suffix_list: list[str], -) -> list[DatasetEntry]: - dataset = [] - for suffix in suffix_list: - for positive_persona, negative_persona in zip( - positive_personas, negative_personas - ): - dataset.append( - DatasetEntry( - positive=template.format(persona=positive_persona) + f" {suffix}", - negative=template.format(persona=negative_persona) + f" {suffix}", - ) - ) - return dataset - @functools.lru_cache(maxsize=1) def load_suffixes() -> list[str]: diff --git a/repeng/utils.py b/repeng/utils.py new file mode 100644 index 0000000..4b7a587 --- /dev/null +++ b/repeng/utils.py @@ -0,0 +1,367 @@ +import json +import copy +import typing +import dataclasses + +import warnings +from loguru import logger + + +@dataclasses.dataclass +class DatasetEntry: + positive: typing.Union[str, typing.List[typing.Dict]] + negative: typing.Union[str, typing.List[typing.Dict]] + + +def make_dataset( + template: typing.Union[str, list], + positive_personas: list[str], + negative_personas: list[str], + suffix_list: typing.Optional[list[str]]=None, +) -> list[DatasetEntry]: + """ + Create a dataset of positive and negative examples based on provided templates and personas. + + Args: + template (Union[str, list]): A string or list of dictionaries containing placeholders {persona} and {suffix}. If a list of dict, they will be automatically turned into a string using the chat template defined in the model repository. If the template is missing the {suffix} placeholder, if template is a str then {suffix} will be appended directly at its end but if template is a list of dict we will crash. + positive_personas (list[str]): A list of positive personas to be used in the dataset. + negative_personas (list[str]): A list of negative personas to be used in the dataset. + suffix_list (Optional: list[str]): A list of suffixes to be used in the dataset. + + Returns: + list[DatasetEntry]: A list of DatasetEntry objects, each containing a positive and negative example. No duplicates are allowed. + + Raises: + ValueError: If the template is neither a string nor a list. + AssertionError: If the template doesn't contain required placeholders, or if there are duplicate items in the dataset. + """ + assert "{persona}" in str(template), "Missing {persona} placeholder in template" + assert len(positive_personas) == len(negative_personas), "You must give the same number of positive and negative personas" + assert len(positive_personas) == len(set(positive_personas)), "Found duplicates in positive personas" + assert len(negative_personas) == len(set(negative_personas)), "Found duplicates in negative personas" + + if suffix_list: + if not "{suffix}" in str(template): + if isinstance(template, str): + template += "{suffix}" + else: + raise Exception("You have to specify a {suffix} placeholder if the template is a dict") + else: + suffix_list = [] + dataset = [] + for suffix in suffix_list: + for positive_persona, negative_persona in zip( + positive_personas, negative_personas + ): + if isinstance(template, str): + positive_template = copy.deepcopy(template).format( + persona=positive_persona, suffix=suffix + ) + negative_template = copy.deepcopy(template).format( + persona=negative_persona, suffix=suffix + ) + + elif isinstance(template, list): + positive_template = copy.deepcopy(template) + for il, l in enumerate(positive_template): + assert isinstance(l, dict), type(l) + for k, v in l.items(): + positive_template[il][k] = v.format( + persona=positive_persona, suffix=suffix + ) + + negative_template = copy.deepcopy(template) + for il, l in enumerate(negative_template): + assert isinstance(l, dict), type(l) + for k, v in l.items(): + negative_template[il][k] = v.format( + persona=negative_persona, suffix=suffix + ) + else: + raise ValueError(type(template)) + + assert positive_template != negative_template, "Error when templating pairs" + dataset.append( + DatasetEntry( + positive=positive_template, + negative=negative_template, + ) + ) + + # check uniqueness + as_strings = [] + for de in dataset: + as_strings.extend([json.dumps(de.positive), json.dumps(de.negative)]) + assert len(as_strings) == len(set(as_strings)), "Found duplicate examples in the dataset" + return dataset + + +def get_model_name(model) -> str: + """ + Retrieve the name of the given model. + + This function attempts to find the name or path of the model by checking + various attributes commonly found in different model implementations. + + Args: + model: The model object to retrieve the name from. + + Returns: + str: The name or path of the model. + + Raises: + ValueError: If the model name cannot be determined. + """ + if hasattr(model, "name_or_path"): + return model.name_or_path + elif hasattr(model, "model"): + return get_model_name(model.model) + elif hasattr(model, "config"): + return model.config.to_dict()["_name_or_path"] + else: + raise ValueError("Couldn't find model name") + +def get_num_hidden_layer(model) -> int: + """ + Retrieve the number of hidden layer in a given model. + + Args: + model: The model object to retrieve the num_hidden_layers from. + + Returns: + int: The number of hidden layers + + Raises: + ValueError: If the model num_hidden_layers cannot be determined. + """ + if hasattr(model.config, "num_hidden_layers"): + num_hidden_layers = model.config.num_hidden_layers + elif hasattr(model.config, "text_config"): + # gemma3 models have a config for text and one for images + num_hidden_layers = model.config.text_config.num_hidden_layers + else: + raise ValueError("Can't find the number of hidden layers") + return num_hidden_layers + + +def autocorrect_chat_templates( + messages: typing.Union[list[list[dict]], list[dict], list[str], str], + tokenizer, + model, + **kwargs, +) -> typing.Union[list[str], str]: + """ + Autocorrect chat templates to ensure compatibility with the given model and tokenizer. + + This function attempts to correct the format of chat messages to match the expected + input format for the specified model and tokenizer. It handles various input types + and applies model-specific corrections when necessary. + + It might be necessary anymore but was needed when huggingface hadn't yet + implemented proper templates. It is nonetheless used as fallback if + model.train crashes because of the template. + + If the tokenizer has no chat template, we crudely turn the input chat + messages into a dialogue. + + Args: + messages (Union[list[list[dict]], list[dict], list[str], str]): The input messages + to be corrected. Can be a single message, a list of messages, or a list of chats. + tokenizer: The tokenizer associated with the model. + model: The model for which the chat templates should be corrected. + kwargs: Any additional kwargs are passed to the tokenizer.__call__ call + + Returns: + Union[list[str], str]: The corrected chat template(s) as a string or list of strings. + + Raises: + ValueError: If the model type is not supported for template correction. + Exception: If some chat messages are still missing after attempting to correct the template. + + Note: + This function includes specific handling for Mistral and LLaMA model variants. + """ + + if isinstance(messages, str): # not a chat template + return messages + elif isinstance(messages, list) and all( + isinstance(mess, list) for mess in messages + ): # list of chats instead of a list of messages + templated = [ + autocorrect_chat_templates(chats, model=model, tokenizer=tokenizer) + for chats in messages + ] + assert all(isinstance(t, str) for t in templated) + assert len(templated) == len( + set(templated) + ), "the dataset should not contain duplicates" + return templated + + # if there is no chat template, make it ourselves + if not tokenizer.chat_template: + output = "" + for m in messages: + role, content = m["role"], m["content"] + output += f"- {role.title()}: {content.strip()}\n" + return output.strip() + + model_name = get_model_name(model).lower() + + assert isinstance(messages, list), "messages should be a list at this point" + assert len(messages), "chat can't be empty" + + for message in messages: + assert isinstance( + message, dict + ), f"messages should be dict, not {type(message)}" + assert sorted(list(message.keys())) == [ + "content", + "role", + ], f"messages should consist of 'content' and 'role' keys only" + assert message["role"] in [ + "user", + "assistant", + "system", + ], f"the role of the message should be user or assistant or system. Found '{ex['role']}'" + assert message[ + "content" + ].strip(), f"message of role '{message['role']}' contains empty string(s)" + + templated = tokenizer.apply_chat_template(messages, tokenize=False, **kwargs) + + if not all(message["content"] in templated for message in messages): + + # see if moving the system message at the end is enough to fix the issue + copied_mes = copy.deepcopy(messages) + sys_message = [mes for mes in copied_mes if mes["role"] == "system"] + assert len(sys_message) == 1, "expected to find a system message" + sys_message = sys_message[0] + copied_mes.remove(sys_message) + copied_mes.append(sys_message) + templated2 = None + try: + templated2 = tokenizer.apply_chat_template( + copied_mes, tokenize=False, **kwargs + ) + except Exception as e: + if ( + not "After the optional system message, conversation roles must alternate user/assistant/user/assistant/..." + in str(e) + ): + raise + if templated2: + if all(message["content"] in templated2 for message in messages): + return templated2 + + copied_mes = copy.deepcopy(messages) + for message in messages: + if message["content"] not in templated: + logger.debug(f"Message '{message['content']}' with role '{message['role']}' is missing after chat template application") + copied_mes = [e for e in copied_mes if e["role"] != "system"] + + first_user_index = [i for i, m in enumerate(copied_mes) if m["role"] == "user"][ + 0 + ] + last_user_index = [i for i, m in enumerate(copied_mes) if m["role"] == "user"][ + -1 + ] + assert ( + copied_mes[0]["role"] == "user" + ), "Expected to find a user message first (or just after the system message)" + + # try to respect the most appropriate chat template + if "mistral" in model_name: + # source: https://github.com/mistralai/cookbook/blob/main/concept-deep-dive/tokenization/templates.md + mistral_versions = { + "v1": [ + "mistral-7b-v0.1", + "mistral-7b-instruct-v0.1", + "mistral-7b-v0.2", + "mistral-7b-instruct-v0.2", + "mixtral-8x7b-v0.1", + "mixtral-8x7b-instruct-v0.1", + "mixtral-8x22b-v0.1", + ], + "v3": [ + "mixtral-8x22b-instruct-v0.1", + "mistral-7b-v0.3", + "mistral-7b-instruct-v0.3", + "codestral-22b-v0.1", + "mathstral-7b-v0.3", + "mamba-codestral-7b-v0.1", + "mistral-large-123b-instruct-2407", + "mistral small 22b instruct 2407", + ], + "v3_tekken": [ + "mistral-nemo-12b-2407", + "mistral-nemo-12b-instruct-2407", + "pixtral-12b-2409", + "ministral-8b-instruct-2410", + "mistral-nemo-instruct-2407", + ], + } + + tokenizer_version = None + for vn, mnames in mistral_versions.items(): + for mname in mnames: + # look for model, including after removing size information like '11b' + if ( + mname in model_name + or "-".join( + [ + mn + for mn in mname.split("-") + if not (mn.endswith("b") and mn.split("b")[0].isdigit()) + ] + ) + in model_name + ): + assert ( + tokenizer_version is None + ), f"Couldn't identify mistral tokenizer version (conflict)" + tokenizer_version = vn + break + assert ( + tokenizer_version is not None + ), f"Couldn't identify mistral tokenizer version (no match found)" + + if tokenizer_version == "v1": + # other source https://huggingface.co/mistralai/Mistral-Nemo-Instruct-2407/discussions/76 + copied_mes[first_user_index][ + "content" + ] = f"{sys_message['content'].rstrip()}\n\n{copied_mes[0]['content'].lstrip()}" + elif tokenizer_version in [ + "v3", + "v3_tekken", + ]: # the difference is mostly about tool handling + copied_mes[last_user_index][ + "content" + ] = f"{sys_message['content'].rstrip()}\n\n{copied_mes[0]['content'].lstrip()}" + else: + raise ValueError(tokenizer_version) + + elif "llama" in model_name: + # according to https://github.com/rohan-paul/LLM-Prompt-Formatting-for-finetuning-Inferencing + copied_mes[first_user_index][ + "content" + ] = f"<>\n{sys_message['content'].rstrip()}\n<>\n\n{copied_mes[first_user_index]['content'].lstrip()}" + else: + # raise ValueError( + # "Besides mistral and llama model, no other chat template correction are implemented" + # ) + warnings.warn("Failed to properly autocorrect the chat template, will use a sane default template") + copied_mes[first_user_index][ + "content" + ] = f"{sys_message['content'].rstrip()}\n\n{copied_mes[first_user_index]['content'].lstrip()}" + + templated = tokenizer.apply_chat_template(copied_mes, tokenize=False, **kwargs) + + if not all(message["content"] in templated for message in messages): + for message in messages: + if message["content"] not in templated: + logger.error(f"Message '{message['content']}' with role '{message['role']}' is STILL missing after chat template application") + raise Exception( + "Some chat messages are still missing after correcting chat template" + ) + + return templated