33import logging
44import numpy as np
55from typing import Dict , List , Optional , Tuple , Any , Union
6- from dataclasses import dataclass
6+ from dataclasses import dataclass , field
77from collections import OrderedDict , defaultdict
88import torch .nn .functional as F
99import torch .nn as nn
@@ -89,6 +89,27 @@ def count_reasoning_tokens(text: str, tokenizer=None) -> int:
8989 MLX_AVAILABLE = False
9090 logger .debug ("MLX framework not available - falling back to PyTorch" )
9191
92+
93+ # Hard ceiling of 4096 by default. Can be lowered via OPTILLM_MAX_TOKENS so a
94+ # single local generation is bounded even when the request (or an approach's
95+ # internal calls) sends no max_tokens -- important for small local models that
96+ # do not reliably emit an EOS token (e.g. the dhara test model), which would
97+ # otherwise ramble up to the full default on every call.
98+ DEFAULT_MAX_NEW_TOKENS = 4096
99+
100+
101+ def _default_max_new_tokens () -> int :
102+ """Default ``max_new_tokens`` for local generation (env-overridable)."""
103+ raw = os .environ .get ("OPTILLM_MAX_TOKENS" )
104+ if raw is None :
105+ return DEFAULT_MAX_NEW_TOKENS
106+ try :
107+ return max (1 , int (raw ))
108+ except (TypeError , ValueError ):
109+ logger .warning ("Ignoring invalid OPTILLM_MAX_TOKENS=%r; using %d" , raw , DEFAULT_MAX_NEW_TOKENS )
110+ return DEFAULT_MAX_NEW_TOKENS
111+
112+
92113@dataclass
93114class ModelConfig :
94115 base_model_id : str
@@ -98,7 +119,7 @@ class ModelConfig:
98119 quantization_bits : int = 4
99120 device_preference : Optional [str ] = None
100121 # Default generation parameters
101- max_new_tokens : int = 4096
122+ max_new_tokens : int = field ( default_factory = _default_max_new_tokens )
102123 do_sample : bool = True
103124 top_p : float = 0.9
104125 top_k : int = 50
@@ -292,7 +313,7 @@ def suggest_mlx_alternative(model_id: str) -> str:
292313class MLXModelConfig :
293314 """Configuration for MLX models"""
294315 model_id : str
295- max_new_tokens : int = 4096
316+ max_new_tokens : int = field ( default_factory = _default_max_new_tokens )
296317 temperature : float = 0.7
297318 top_p : float = 0.9
298319 repetition_penalty : float = 1.0
@@ -1268,16 +1289,46 @@ def setup_tokenizer(self, tokenizer: AutoTokenizer) -> AutoTokenizer:
12681289
12691290 return tokenizer
12701291
1292+ def _resolve_eos_token_ids (self ):
1293+ """Resolve the effective end-of-sequence token id(s) for generation.
1294+
1295+ Prefer the model's own ``generation_config.eos_token_id``. Chat models
1296+ commonly set it to the chat-turn end token (e.g. ``<|im_end|>``), which
1297+ can differ from the tokenizer's ``eos_token_id`` (often the base-model
1298+ ``<|end_of_text|>``). Passing only the tokenizer eos to ``generate`` there
1299+ means the model never stops on its real turn-end token and rambles up to
1300+ ``max_new_tokens`` -- e.g. dhara-250m's ChatML ends at ``<|im_end|>`` but
1301+ its tokenizer eos is ``<|end_of_text|>``.
1302+
1303+ The tokenizer eos is merged in as a fallback so a model that only emits
1304+ the base eos still terminates. Returns an int, a list of ints, or None.
1305+ """
1306+ ids : List [int ] = []
1307+ gen_cfg = getattr (self .current_model , "generation_config" , None )
1308+ gc_eos = getattr (gen_cfg , "eos_token_id" , None ) if gen_cfg is not None else None
1309+ if isinstance (gc_eos , int ):
1310+ ids .append (gc_eos )
1311+ elif isinstance (gc_eos , (list , tuple )):
1312+ ids .extend (int (x ) for x in gc_eos if isinstance (x , int ))
1313+ tok_eos = self .tokenizer .eos_token_id
1314+ if isinstance (tok_eos , int ):
1315+ ids .append (tok_eos )
1316+ seen = set ()
1317+ resolved = [x for x in ids if not (x in seen or seen .add (x ))]
1318+ if not resolved :
1319+ return None
1320+ return resolved [0 ] if len (resolved ) == 1 else resolved
1321+
12711322 def get_optimized_generation_config (self , generation_params : Optional [Dict [str , Any ]] = None ) -> Dict :
12721323 """Get optimized generation config"""
12731324 config = {
1274- "max_new_tokens" : generation_params .get ("max_new_tokens" , 4096 ),
1325+ "max_new_tokens" : generation_params .get ("max_new_tokens" , _default_max_new_tokens () ),
12751326 "do_sample" : generation_params .get ("temperature" , 1.0 ) > 0 ,
12761327 "temperature" : generation_params .get ("temperature" , 1.0 ),
12771328 "top_p" : generation_params .get ("top_p" , 0.95 ),
12781329 "num_return_sequences" : generation_params .get ("num_return_sequences" , 1 ),
12791330 "pad_token_id" : self .tokenizer .pad_token_id ,
1280- "eos_token_id" : self .tokenizer . eos_token_id ,
1331+ "eos_token_id" : self ._resolve_eos_token_ids () ,
12811332 "return_dict_in_generate" : True ,
12821333 "output_scores" : generation_params .get ("logprobs" , False ),
12831334 "use_cache" : True
@@ -1571,13 +1622,13 @@ def process_batch(
15711622 if batch_prompts : # If there are any uncached prompts
15721623 # Configure generation parameters
15731624 base_params = {
1574- "max_new_tokens" : generation_params .get ("max_new_tokens" , 4096 ) if generation_params else self .model_config .max_new_tokens ,
1625+ "max_new_tokens" : generation_params .get ("max_new_tokens" , _default_max_new_tokens () ) if generation_params else self .model_config .max_new_tokens ,
15751626 "do_sample" : generation_params .get ("temperature" , 1.0 ) > 0 if generation_params else self .model_config .do_sample ,
15761627 "temperature" : generation_params .get ("temperature" , 1.0 ) if generation_params else self .model_config .temperature ,
15771628 "top_p" : generation_params .get ("top_p" , 1.0 ) if generation_params else self .model_config .top_p ,
15781629 "num_return_sequences" : n ,
15791630 "pad_token_id" : self .tokenizer .pad_token_id ,
1580- "eos_token_id" : self .tokenizer . eos_token_id ,
1631+ "eos_token_id" : self ._resolve_eos_token_ids () ,
15811632 }
15821633
15831634 # Add optional parameters if specified
@@ -1900,7 +1951,7 @@ def create(
19001951
19011952 # Use directly available parameters for entropy decoding
19021953 entropy_params = {
1903- "max_new_tokens" : max_tokens if max_tokens is not None else 4096 ,
1954+ "max_new_tokens" : max_tokens if max_tokens is not None else _default_max_new_tokens () ,
19041955 "temperature" : temperature ,
19051956 "top_p" : top_p ,
19061957 "top_k" : top_k ,
@@ -2046,7 +2097,7 @@ def create(
20462097 "temperature" : temperature ,
20472098 "top_p" : top_p ,
20482099 "num_return_sequences" : n ,
2049- "max_new_tokens" : max_tokens if max_tokens is not None else 4096 ,
2100+ "max_new_tokens" : max_tokens if max_tokens is not None else _default_max_new_tokens () ,
20502101 "presence_penalty" : presence_penalty ,
20512102 "frequency_penalty" : frequency_penalty ,
20522103 "stop_sequences" : [stop ] if isinstance (stop , str ) else stop ,
0 commit comments