1717)
1818
1919from backend .config import set_config_dir , global_state , config_filename
20- from backend .models import get_loaded_model
20+ from backend .models import set_model_loaded_callback
2121from backend .prompts import prompt_formats
2222from backend .util import MultiTimer
23+ import backend .models as models # Import as module to avoid circular dependency
2324import threading
2425
2526session_list : dict or None = None
2627current_session = None
2728
29+ def handle_model_loaded (model ):
30+ """Handle model loading - only update new sessions with model params"""
31+ pass
32+
33+ # Register callback to handle model loading
34+ set_model_loaded_callback (handle_model_loaded )
35+
2836# Cancel
2937
3038abort_event = threading .Event ()
@@ -92,9 +100,14 @@ def delete_session(d_session):
92100 current_session = None
93101
94102
95- def get_default_session_settings ():
96- return \
97- {
103+ def get_default_session_settings (use_model_params = False ):
104+ """Get default session settings
105+
106+ Args:
107+ use_model_params: If True and a model is loaded with custom params,
108+ apply those params instead of defaults
109+ """
110+ settings = {
98111 "prompt_format" : "Chat-RP" ,
99112 "roles" : [ "User" , "Assistant" , "" , "" , "" , "" , "" , "" ],
100113 "system_prompt_default" : True ,
@@ -119,6 +132,23 @@ def get_default_session_settings():
119132 "temperature_last" : False ,
120133 "skew" : 0.0 ,
121134 }
135+
136+ if use_model_params :
137+ # If requested, try to use model parameters
138+ loaded_model = models .get_loaded_model ()
139+ if loaded_model is not None :
140+ model_dict = loaded_model .model_dict
141+ # Only apply if model has custom params defined
142+ if any (param in model_dict for param in ["temperature" , "top_k" , "top_p" , "repp" ]):
143+ settings .update ({
144+ "temperature" : model_dict .get ("temperature" , settings ["temperature" ]),
145+ "top_k" : model_dict .get ("top_k" , settings ["top_k" ]),
146+ "top_p" : model_dict .get ("top_p" , settings ["top_p" ]),
147+ "repp" : model_dict .get ("repp" , settings ["repp" ])
148+ })
149+ print ("Updated settings with model params:" , settings )
150+
151+ return settings
122152
123153class Session :
124154
@@ -145,7 +175,8 @@ def init_new(self):
145175 self .session_uuid = str (uuid .uuid4 ())
146176 self .history = []
147177 # self.mode = ""
148- self .settings = get_default_session_settings ()
178+ # New sessions get app defaults
179+ self .settings = get_default_session_settings (use_model_params = False )
149180
150181
151182 def to_json (self ):
@@ -163,9 +194,13 @@ def from_json(self, j):
163194 self .session_uuid = j ["session_uuid" ]
164195 self .history = j ["history" ]
165196 # self.mode = j["mode"]
166- settings = get_default_session_settings ()
167- if "settings" in j : settings .update (j ["settings" ])
168- self .settings = settings
197+
198+ # Start with hardcoded defaults (no model params)
199+ self .settings = get_default_session_settings (use_model_params = False )
200+
201+ # Apply ALL saved settings including sampling params
202+ if "settings" in j :
203+ self .settings .update (j ["settings" ])
169204
170205
171206 def load (self ):
@@ -244,7 +279,7 @@ def create_context(self, prompt_format, max_len, min_len, uptoblock = None, pref
244279
245280 def create_context_instruct (self , prompt_format , max_len , min_len , uptoblock = None , prefix = "" ):
246281
247- tokenizer = get_loaded_model ().tokenizer
282+ tokenizer = models . get_loaded_model ().tokenizer
248283 prompts = []
249284 responses = []
250285
@@ -347,7 +382,7 @@ def create_context_instruct(self, prompt_format, max_len, min_len, uptoblock = N
347382
348383 def create_context_raw (self , prompt_format , max_len , min_len , uptoblock = None , prefix = "" ):
349384
350- tokenizer = get_loaded_model ().tokenizer
385+ tokenizer = models . get_loaded_model ().tokenizer
351386 history_copy = []
352387 for h in self .history :
353388 if h ["block_uuid" ] == uptoblock : break
@@ -413,16 +448,17 @@ def generate(self, data):
413448 gen_prefix = data .get ("prefix" , "" )
414449 block_id = data .get ("block_id" , None )
415450
416- if get_loaded_model () is None :
451+ if models . get_loaded_model () is None :
417452 packet = { "result" : "fail" , "error" : "No model loaded." }
418453 yield json .dumps (packet ) + "\n "
419454 return packet
420455
421- model = get_loaded_model ().model
422- generator = get_loaded_model ().generator
423- tokenizer = get_loaded_model ().tokenizer
424- cache = get_loaded_model ().cache
425- speculative_mode = get_loaded_model ().speculative_mode
456+ loaded_model = models .get_loaded_model ()
457+ model = loaded_model .model
458+ generator = loaded_model .generator
459+ tokenizer = loaded_model .tokenizer
460+ cache = loaded_model .cache
461+ speculative_mode = loaded_model .speculative_mode
426462
427463 prompt_format = prompt_formats [self .settings ["prompt_format" ]]()
428464
0 commit comments