Fixing compare feature
This commit is contained in:
parent
b6799a21a5
commit
fdeccec259
1 changed files with 89 additions and 69 deletions
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue