Skip to content
3 changes: 2 additions & 1 deletion mellea/backends/huggingface.py
Original file line number Diff line number Diff line change
Expand Up @@ -1283,7 +1283,8 @@ async def _generate_from_intrinsic(

# Two load-bearing assumptions: (1) __getattr__ falls through for attributes
# accessed by generate_with_transformers (model.device, model.vocab_size,
# model.generation_config). (2) chat_completion_request_to_transformers_inputs
# model.generation_config, and for aLoRA activation model.peft_config,
# model.active_adapters, model.modules()). (2) chat_completion_request_to_transformers_inputs
# always sets return_dict_in_generate=True, so .generate() always returns a
# GenerateDecoderOnlyOutput — if that ever changes, the cell stays None and
# logits silently won't be populated.
Expand Down
164 changes: 163 additions & 1 deletion mellea/formatters/granite/base/util.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,10 +6,12 @@
# Standard
from __future__ import annotations

import contextlib
import itertools
import json
import os
import uuid
from collections.abc import Iterator, Mapping
from typing import TYPE_CHECKING, Any, cast

# Third Party
Expand Down Expand Up @@ -378,6 +380,160 @@ def chat_completion_request_to_transformers_inputs(
return generate_input, other_input


# aLoRA adapter names already warned about for a missing invocation sequence,
# so a mispackaged adapter logs once per process rather than once per call.
_ALORA_MISSING_INVOCATION_WARNED: set[str] = set()


def _generation_setting(
Comment thread
planetf1 marked this conversation as resolved.
model: PreTrainedModel, generate_kwargs: Mapping[str, Any], name: str
) -> int:
"""Effective integer `generate()` setting, resolved in `generate()`'s order.

A flat kwarg wins, then a `generation_config=` kwarg, then the model's own
`generation_config`; 1 if none sets it. Unset `GenerationConfig` fields
are `None`, so each level falls through to the next.
"""
value = generate_kwargs.get(name)
if value is None:
value = getattr(generate_kwargs.get("generation_config"), name, None)
if value is None:
value = getattr(getattr(model, "generation_config", None), name, None)
return value if isinstance(value, int) else 1


@contextlib.contextmanager
def _alora_activation_context(
model: PreTrainedModel,
input_tokens: torch.Tensor | None,
generate_kwargs: Mapping[str, Any] | None = None,
) -> Iterator[None]:
"""Context manager that makes aLoRA weights actually apply during `generate()`.

PEFT computes the per-token aLoRA activation mask (`alora_offsets`) only
from inside its `PeftModel` wrapper's own `generate()`/`forward()`
overrides. A model that loads adapters via transformers' native
`PeftAdapterMixin.load_adapter()` (as `LocalHFBackend` does) is never
wrapped in `PeftModel`, so without this the aLoRA weights silently never
apply and generation degrades to base-model behaviour. This mirrors what
`PeftModel` does internally (`_enable_peft_forward_hooks` in
`peft/tuners/lora/model.py`): compute the offsets with PEFT's
`calculate_alora_offsets` helper (internal to
`peft.tuners.lora.variants`, not part of PEFT's exported API), then inject
them into every LoRA layer's forward via temporary pre-forward hooks. The
hook is required because transformers' model code does not generally
propagate `**kwargs` down to the projection `nn.Linear` calls, so passing
`alora_offsets` to `generate()` alone is not enough.

Yields immediately (no-op) unless exactly one adapter is active and its
PEFT config declares a non-empty `alora_invocation_tokens`. When the
declared sequence does not occur in `input_tokens`, the adapter cannot
activate: no hooks are registered (so generation uses the base-model
weights, as PEFT would) and a WARNING naming the adapter is logged once
per process. Offsets are repeated per returned sequence, since
`generate()` expands the prompt batch for `num_return_sequences`. Hooks
are removed on exit, including on error. Safe to call on any model; the
PEFT imports happen lazily so this module keeps working without the `hf`
extra.

Args:
model: Hugging Face model object (or a proxy that forwards attribute
access to one).
input_tokens: Prompt token-id tensor on the model's device, as
passed to `generate(input_ids=...)`.
generate_kwargs: The other keyword arguments the caller is about to
pass to `generate()`, read for `num_beams` and
`num_return_sequences`.

Yields:
None. Active for the duration of the `with` block.

Raises:
ValueError: Beam search (`num_beams > 1`) was requested while an aLoRA
adapter can activate; PEFT does not support that combination.
"""
peft_config = getattr(model, "peft_config", None)
if not peft_config:
yield
return
try:
active_adapters = model.active_adapters
# `active_adapters` is a method on transformers' native
# PeftAdapterMixin but a property on peft's PeftModel wrapper.
active = cast(
"list[str]",
active_adapters() if callable(active_adapters) else active_adapters,
)
except ValueError as e:
# transformers raises
# `ValueError("No adapter loaded. Please load an adapter first.")`
# when nothing is loaded; anything else is a real fault.
if "No adapter loaded" not in str(e):
raise
active = []
if len(active) != 1:
yield
return
adapter_name = active[0]
invocation = getattr(peft_config.get(adapter_name), "alora_invocation_tokens", None)
if not invocation:
yield
return
if input_tokens is None or not hasattr(input_tokens, "device"):
yield
return

# Third Party (lazy: peft is an optional `hf`-extra dependency)
from peft.tuners.lora import (
layer as _peft_lora_layer,
variants as _peft_lora_variants,
)

offsets = _peft_lora_variants.calculate_alora_offsets(
peft_config, adapter_name, input_tokens
)
if all(offset is None for offset in offsets):
if adapter_name not in _ALORA_MISSING_INVOCATION_WARNED:
_ALORA_MISSING_INVOCATION_WARNED.add(adapter_name)
MelleaLogger.get_logger().warning(
f"aLoRA adapter {adapter_name!r} is active but its declared "
f"invocation sequence {list(invocation)} does not occur in the "
"prompt, so generation used base-model weights only. Check that "
"the adapter's io.yaml instruction or chat template produces "
"that token sequence."
)
yield
return

generate_kwargs = generate_kwargs or {}
if _generation_setting(model, generate_kwargs, "num_beams") > 1:
# Same refusal as PeftModel's own path (`_enable_peft_forward_hooks`).
raise ValueError(
f"Beam search is not supported for aLoRA adapters (adapter {adapter_name!r})."
)
# generate() expands each prompt row `num_return_sequences` times via
# `repeat_interleave`, so repeat each row's offset the same way.
n_sequences = _generation_setting(model, generate_kwargs, "num_return_sequences")
offsets = [offset for offset in offsets for _ in range(n_sequences)]

def _inject_offsets(
module: torch.nn.Module, args: tuple[Any, ...], kwargs: dict[str, Any]
) -> tuple[tuple[Any, ...], dict[str, Any]]:
kwargs["alora_offsets"] = offsets
return args, kwargs

handles = [
module.register_forward_pre_hook(_inject_offsets, with_kwargs=True)
for module in model.modules()
if isinstance(module, _peft_lora_layer.LoraLayer)
]
try:
yield
finally:
for handle in handles:
handle.remove()


def generate_with_transformers(
tokenizer: PreTrainedTokenizerBase,
model: PreTrainedModel,
Expand All @@ -403,6 +559,11 @@ def generate_with_transformers(

Returns:
A chat completion response in OpenAI format.

Raises:
ValueError: Beam search (`num_beams > 1`) was requested while an aLoRA
adapter is active and its invocation sequence is in the prompt;
PEFT does not support that combination.
"""
with import_optional("torch"):
# Third Party
Expand All @@ -412,7 +573,8 @@ def generate_with_transformers(
generate_input = generate_input.copy()
del generate_input["input_tokens"]

generate_result = model.generate(input_ids=input_tokens, **generate_input) # type: ignore[operator]
with _alora_activation_context(model, input_tokens, generate_input):
generate_result = model.generate(input_ids=input_tokens, **generate_input) # type: ignore[operator]

# Result is a a 2D tensor of shape (num responses, prompt + max generated tokens)
# containing tokens, plus a tuple of <max generated tokens> tensors of shape
Expand Down
Loading
Loading