Fixing compare feature

This commit is contained in:
Manan17 2026-02-19 20:15:44 +00:00
commit fdeccec259

View file

@ -38,9 +38,13 @@ class InferenceBackend:
]
self.device = get_device().value
# Thread safety
# Thread safety — _generation_lock serializes model.generate() calls.
# Must be a regular Lock (NOT RLock) because in async FastAPI, multiple
# requests share the same event-loop thread, so RLock reentrancy lets
# concurrent compare-mode requests race on the GPU. The lock is
# acquired by the *background generation thread*, not the event-loop.
import threading
self._generation_lock = threading.RLock()
self._generation_lock = threading.Lock()
self._model_state_lock = threading.Lock()
logger.info(f"InferenceBackend initialized on {self.device}")
@ -448,9 +452,10 @@ class InferenceBackend:
"""
Apply adapter state before generation. Must be called under _generation_lock.
Uses revert_to_base_model() / activate_lora_adapter() which work correctly
for models loaded by Unsloth as complete PeftModels (via model.unload() /
model.load_adapter()), matching the proven pattern from the Gradio eval page.
Uses PEFT's disable_adapter_layers() / enable_adapter_layers() which toggle
a boolean flag on each LoRA layer. Unsloth's fast_linear_forward checks this
flag (proj.disable_adapters) and skips LoRA computation when True.
This is non-destructive no model unloading/reloading needed.
Args:
use_adapter: None = no change, False = disable (base model),
@ -464,32 +469,34 @@ class InferenceBackend:
return
model_info = self.models[base]
model = model_info.get("model")
if model is None:
return
if use_adapter is False:
# Revert to pure base model by unloading adapter weights
logger.info(f"Compare mode: reverting '{base}' to base model for generation")
self.revert_to_base_model(base)
# Disable LoRA layers → base model output
if isinstance(model, (PeftModel, PeftModelForCausalLM)):
logger.info(f"Compare mode: disabling adapters on '{base}' for base model generation")
model.base_model.disable_adapter_layers()
else:
logger.info(f"Compare mode: model '{base}' is not a PeftModel, already base")
elif use_adapter is True:
# Activate the LoRA adapter from the original model path
lora_path = model_info.get("model_path")
if lora_path and model_info.get("is_lora"):
logger.info(f"Compare mode: activating LoRA adapter from '{lora_path}' on '{base}'")
self.activate_lora_adapter(base, lora_path)
# Re-enable LoRA layers → adapter output
if isinstance(model, (PeftModel, PeftModelForCausalLM)):
logger.info(f"Compare mode: enabling adapters on '{base}' for LoRA generation")
model.base_model.enable_adapter_layers()
else:
# Fallback for dynamically attached adapters
loaded = model_info.get("loaded_adapters", {})
if loaded:
adapter_name = list(loaded.keys())[-1]
logger.info(f"Compare mode: enabling adapter '{adapter_name}' on '{base}'")
self.set_active_adapter(base, adapter_name)
else:
logger.warning("use_adapter=true but no adapter path/adapters on model")
logger.warning("use_adapter=true but model is not a PeftModel")
elif isinstance(use_adapter, str):
# Activate a specific adapter by path
logger.info(f"Compare mode: activating specific adapter '{use_adapter}' on '{base}'")
self.activate_lora_adapter(base, use_adapter)
# Enable adapters and set the specific one active
if isinstance(model, (PeftModel, PeftModelForCausalLM)):
logger.info(f"Compare mode: enabling adapter '{use_adapter}' on '{base}'")
model.base_model.enable_adapter_layers()
self.set_active_adapter(base, use_adapter)
else:
logger.warning(f"use_adapter='{use_adapter}' but model is not a PeftModel")
def generate_with_adapter_control(
self,
@ -500,18 +507,18 @@ class InferenceBackend:
"""
Thread-safe generation with optional adapter toggling.
Acquires the generation lock, applies adapter state, then generates.
This ensures adapter toggle + generation are atomic critical for
compare mode where base and LoRA panes fire concurrently.
The adapter toggle + model.generate() are serialized by _generation_lock
inside the background generation thread NOT in the event-loop thread.
This prevents the RLock-reentrant race that occurs when two async SSE
handlers share the same event-loop thread.
Args:
use_adapter: Adapter control (None/False/True/str). See _apply_adapter_state.
**gen_kwargs: Forwarded to generate_chat_response.
"""
with self._generation_lock:
self._apply_adapter_state(use_adapter)
# Delegate to the lock-free generation path
yield from self._generate_chat_response_inner(cancel_event=cancel_event, **gen_kwargs)
yield from self._generate_chat_response_inner(
cancel_event=cancel_event, _adapter_state=use_adapter, **gen_kwargs
)
def generate_chat_response(self,
messages: list,
@ -526,22 +533,20 @@ class InferenceBackend:
cancel_event=None) -> Generator[str, None, None]:
"""
Generate response for text or vision models.
Acquires the generation lock. For adapter-controlled generation,
use generate_with_adapter_control() instead.
The generation lock is acquired by the background generation thread.
"""
with self._generation_lock:
yield from self._generate_chat_response_inner(
messages=messages,
system_prompt=system_prompt,
image=image,
temperature=temperature,
top_p=top_p,
top_k=top_k,
min_p=min_p,
max_new_tokens=max_new_tokens,
repetition_penalty=repetition_penalty,
cancel_event=cancel_event,
)
yield from self._generate_chat_response_inner(
messages=messages,
system_prompt=system_prompt,
image=image,
temperature=temperature,
top_p=top_p,
top_k=top_k,
min_p=min_p,
max_new_tokens=max_new_tokens,
repetition_penalty=repetition_penalty,
cancel_event=cancel_event,
)
def _generate_chat_response_inner(self,
messages: list,
@ -553,10 +558,14 @@ class InferenceBackend:
min_p: float = 0.0,
max_new_tokens: int = 256,
repetition_penalty: float = 1.1,
cancel_event=None) -> Generator[str, None, None]:
cancel_event=None,
_adapter_state=None) -> Generator[str, None, None]:
"""
Inner generation logic (no lock). Called by both generate_chat_response
Inner generation logic. Called by both generate_chat_response
and generate_with_adapter_control.
_adapter_state is passed to generate_stream/vision so the background
thread can toggle adapters under the generation lock.
"""
if not self.active_model_name:
yield "Error: No active model"
@ -616,6 +625,7 @@ class InferenceBackend:
yield from self.generate_stream(
formatted_prompt, temperature, top_p, top_k, min_p, max_new_tokens, repetition_penalty,
cancel_event=cancel_event,
_adapter_state=_adapter_state,
)
def _generate_vision_response(self, messages, system_prompt, image,
@ -677,6 +687,7 @@ class InferenceBackend:
streamer=streamer,
max_new_tokens=max_new_tokens,
use_cache=True,
do_sample=temperature > 0,
temperature=temperature,
top_p=top_p,
top_k=top_k,
@ -686,16 +697,17 @@ class InferenceBackend:
err: dict[str, str] = {}
def generate_fn():
try:
model.generate(**generation_kwargs)
except Exception as e:
err["msg"] = str(e)
logger.error(f"Vision generation error in thread: {e}")
finally:
with self._generation_lock:
try:
streamer.end()
except Exception:
pass
model.generate(**generation_kwargs)
except Exception as e:
err["msg"] = str(e)
logger.error(f"Vision generation error in thread: {e}")
finally:
try:
streamer.end()
except Exception:
pass
thread = threading.Thread(target=generate_fn)
thread.start()
@ -741,8 +753,13 @@ class InferenceBackend:
min_p: float = 0.0,
max_new_tokens: int = 256,
repetition_penalty: float = 1.1,
cancel_event=None) -> Generator[str, None, None]:
"""Generate streaming text response (text models only)."""
cancel_event=None,
_adapter_state=None) -> Generator[str, None, None]:
"""Generate streaming text response (text models only).
_adapter_state: if not None, the background thread toggles adapters
before model.generate(), all under _generation_lock.
"""
if not self.active_model_name:
yield "Error: No active model"
return
@ -773,7 +790,7 @@ class InferenceBackend:
top_k=top_k,
min_p=min_p,
repetition_penalty=repetition_penalty,
do_sample=True,
do_sample=temperature > 0,
eos_token_id=tokenizer.eos_token_id,
pad_token_id=tokenizer.eos_token_id if tokenizer.pad_token_id is None else tokenizer.pad_token_id,
)
@ -795,16 +812,19 @@ class InferenceBackend:
)
def generate_fn():
try:
model.generate(**generation_kwargs)
except Exception as e:
err["msg"] = str(e)
logger.error(f"Generation error: {e}")
finally:
with self._generation_lock:
try:
streamer.end()
except Exception:
pass
if _adapter_state is not None:
self._apply_adapter_state(_adapter_state)
model.generate(**generation_kwargs)
except Exception as e:
err["msg"] = str(e)
logger.error(f"Generation error: {e}")
finally:
try:
streamer.end()
except Exception:
pass
err: dict[str, str] = {}
thread = threading.Thread(target=generate_fn)