Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
51 changes: 51 additions & 0 deletions align_system/algorithms/argmax_alignment_adm_component.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,51 @@
from __future__ import annotations

from align_system.algorithms.abstracts import ADMComponent
from align_system.utils import logging

log = logging.getLogger(__name__)


class ArgmaxAlignmentADMComponent(ADMComponent):
"""
Alignment step that picks the choice with the highest predicted
KDMA score, averaged across all attributes and samples.

Replaces the ADEPT random effects model for domains (like AI2Thor)
where no calibrated statistical model exists. For single-attribute
pipelines this reduces to a plain argmax over the LLM's scores.
"""

def __init__(self, attributes=None):
self.attributes = attributes or {}

def run_returns(self):
return ("chosen_choice", "best_sample_idx", "alignment_info")

def run(self, attribute_prediction_scores, alignment_target=None):
"""
attribute_prediction_scores: dict[choice_str, dict[kdma, list[float]]]
"""
choice_totals: dict[str, float] = {}

for choice, attr_scores in attribute_prediction_scores.items():
total = 0.0
count = 0
for kdma, scores in attr_scores.items():
vals = scores if isinstance(scores, list) else [scores]
if vals:
total += sum(vals) / len(vals)
count += 1
choice_totals[choice] = total / count if count else 0.0

best_choice = max(choice_totals, key=choice_totals.get)

log.info(f"[ArgmaxAlignment] scores: {choice_totals}")
log.info(f"[ArgmaxAlignment] chosen: {best_choice}")

alignment_info = {
"source": type(self).__name__,
"choice_scores": choice_totals,
}

return best_choice, 0, alignment_info
6 changes: 3 additions & 3 deletions align_system/algorithms/misc_itm_adm_components.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,8 +15,8 @@ def run_returns(self):
return ('chosen_action')

def run(self,
choices,
actions,
choices=None,
chosen_choice=None,
chosen_action=None,
justification=None):
Expand Down Expand Up @@ -76,8 +76,8 @@ def run_returns(self):
return 'choice_info'

def run(self,
choices,
actions,
choices=None,
alignment_target=None,
attribute_prediction_scores=None,
attribute_relevance=None,
Expand Down Expand Up @@ -105,7 +105,7 @@ def run(self,

true_kdma_values = {}
true_relevance = {}
for choice, action in zip(choices, actions):
for choice, action in zip(choices or [], actions):
if action.kdma_association is not None:
true_kdma_values[choice] = action.kdma_association
for kdma in target_kdmas:
Expand Down
129 changes: 129 additions & 0 deletions align_system/algorithms/ollama_inference_engine.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,129 @@
from __future__ import annotations

import json

import ollama

from align_system.algorithms.abstracts import StructuredInferenceEngine
from align_system.utils import logging

log = logging.getLogger(__name__)


class OllamaInferenceEngine(StructuredInferenceEngine):
"""
StructuredInferenceEngine backed by a local Ollama model.

Uses Ollama's native structured output support
(https://ollama.com/blog/structured-outputs): the JSON schema is
passed as the `format` parameter so the server constrains the
output to match the schema.
"""

def __init__(
self,
model: str = "gemma4:12b",
temperature: float = 0.0,
num_ctx: int = 8192,
num_predict: int = 4096,
max_retries: int = 2,
):
self.model = model
self.temperature = temperature
self.num_ctx = num_ctx
self.num_predict = num_predict
self.max_retries = max_retries

def dialog_to_prompt(self, dialog) -> str:
"""
Flatten a dialog list into a plain-text prompt for Ollama.

System messages are prepended as an unlabelled block so they
land at the top; user/assistant turns follow with role labels.
"""
system_parts = []
turn_parts = []

for elem in dialog:
role = elem.role if hasattr(elem, "role") else elem["role"]
content = elem.content if hasattr(elem, "content") else elem["content"]

if role == "system":
system_parts.append(content)
else:
turn_parts.append(f"[{role.upper()}]\n{content}")

parts = []
if system_parts:
parts.append("\n\n".join(system_parts))
parts.extend(turn_parts)
return "\n\n".join(parts)

def run_inference(self, prompts, schema: str, temperature: float = None):
"""
Run inference for each prompt string and return parsed JSON dicts.

`schema` is a JSON Schema string passed to Ollama's `format`
parameter for server-side constrained generation.
"""
format_schema = json.loads(schema)
effective_temperature = self.temperature if temperature is None else temperature

if single_prompt := isinstance(prompts, str):
prompts = [prompts]

results = []
for prompt in prompts:
for attempt in range(self.max_retries + 1):
# On retries, sample with some temperature so a greedy
# engine doesn't just reproduce the same bad output
retry_temperature = (effective_temperature if attempt == 0
else max(effective_temperature, 0.2))
resp = ollama.generate(
model=self.model,
prompt=prompt,
format=format_schema,
options={"temperature": retry_temperature,
"num_ctx": self.num_ctx,
"num_predict": self.num_predict},
)
text = resp["response"]
log.debug(f"[OllamaInferenceEngine] raw response:\n{text}")

try:
results.append(self._parse_json_response(text))
break
except (json.JSONDecodeError, RuntimeError) as e:
if attempt == self.max_retries:
raise
log.warning(f"[OllamaInferenceEngine] failed to parse "
f"response (attempt {attempt + 1} of "
f"{self.max_retries + 1}): {e}; retrying")

return results[0] if single_prompt else results

@staticmethod
def _parse_json_response(text: str):
"""
Parse the first JSON value in the response, tolerating trailing
garbage (some models emit extra text after the schema-constrained
JSON despite the `format` parameter).
"""
stripped = text.strip()
if not stripped:
raise RuntimeError(
"Ollama returned an empty response; the model may not "
"support structured output via the `format` parameter")

obj, end = json.JSONDecoder().raw_decode(stripped)
trailing = stripped[end:].strip()
if trailing:
log.warning(f"[OllamaInferenceEngine] ignoring trailing data "
f"after JSON response: {trailing[:100]!r}")
return obj

def cache_repr(self) -> str:
return (
f"OllamaInferenceEngine(model={self.model}, "
f"temperature={self.temperature}, num_ctx={self.num_ctx})"
)
46 changes: 34 additions & 12 deletions align_system/algorithms/outlines_inference_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -73,6 +73,11 @@ def __init__(
# newer verion of outlines fixes this issue, but we are blocked with the vllm dependency
self.model.tokenizer.is_llama = True

# If generation_kwargs includes temperature, enable sampling in the model's
# generation_config so transformers doesn't warn that temperature is invalid.
if self.generation_kwargs.get("temperature", 0.0) > 0:
self.model.model.generation_config.do_sample = True

def dialog_to_prompt(self, dialog):
tokenizer = self.model.tokenizer.tokenizer

Expand Down Expand Up @@ -127,40 +132,57 @@ def run_in_batches(
outputs.extend(output)
return outputs

def run_inference(self, prompts, schema):
def run_inference(self, prompts, schema, temperature: float = None):
json_schema = JsonSchema(schema, whitespace_pattern=r"[ ]?")

generator = outlines.Generator(self.model, json_schema)

gen_kwargs = dict(self.generation_kwargs)
if temperature is not None:
gen_kwargs["temperature"] = temperature
gen_kwargs["do_sample"] = True

if isinstance(prompts, str):
output = generator(
prompts,
max_new_tokens=self.max_generator_tokens,
**self.generation_kwargs,
**gen_kwargs,
)
return json.loads(output)
try:
return json.loads(output)
except json.JSONDecodeError as e:
raise RuntimeError(
f"Failed to parse structured generation output as JSON "
f"(output may be truncated; consider increasing "
f"max_generator_tokens above {self.max_generator_tokens}). "
f"Raw output: {output!r}. Original error: {e}"
) from e
elif isinstance(prompts, Iterable):
output = self.run_in_batches(
generator.batch,
prompts,
self.inference_batch_size,
self.max_generator_tokens,
**self.generation_kwargs,
**gen_kwargs,
)
return [json.loads(r) for r in output]
try:
return [json.loads(r) for r in output]
except json.JSONDecodeError as e:
raise RuntimeError(
f"Failed to parse structured generation output as JSON "
f"(output may be truncated; consider increasing "
f"max_generator_tokens above {self.max_generator_tokens}). "
f"Raw output: {output!r}. Original error: {e}"
) from e
else:
raise TypeError(
"Don't know how to run inference on provided `prompts` object"
)

def run_inference_unstructured(self, prompts):
generator = outlines.generate.regex(
self.model,
r".*", # "allow anything" regex
**self.generation_kwargs,
)
generator = outlines.Generator(self.model)

if isinstance(prompts, str):
return generator(prompts, self.max_generator_tokens)
return generator(prompts, max_new_tokens=self.max_generator_tokens, **self.generation_kwargs)
elif isinstance(prompts, Iterable):
return self.run_in_batches(
generator, prompts, self.inference_batch_size, self.max_generator_tokens
Expand Down
30 changes: 27 additions & 3 deletions align_system/algorithms/pipeline_adm.py

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think it's fine to have the history tracking as you have it in here for now, but I'm more inclined to merge Yoni's approach on this: https://github.com/ITM-Kitware/align-system/pull/277/changes#diff-ea512e45fac46d4935ce85a4837bdbcd27b5891a09c1ac5f6aad038076f4d497

As it maintains the full working_output history.

Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
from collections import deque
from collections.abc import Iterable
from timeit import default_timer as timer

Expand All @@ -8,8 +9,9 @@


class PipelineADM(ActionBasedADM):
def __init__(self, steps: list[ADMComponent]):
def __init__(self, steps: list[ADMComponent], history_window=None):
self.steps = steps
self.history = deque(maxlen=history_window)

def choose_action(self,
scenario_state,
Expand All @@ -28,8 +30,9 @@ def choose_action(self,
step_returns = step.run_returns()

start_time = timer()
# Run the step
run_output = call_with_coerced_args(step.run, working_output)
# Run the step, temporarily adding historical working outputs to working_output
working_output_with_history = {**{"history": list(self.history)}, **working_output}
run_output = call_with_coerced_args(step.run, working_output_with_history)
end_time = timer()

per_step_timing_stats.append(
Expand Down Expand Up @@ -74,4 +77,25 @@ def choose_action(self,
working_output.setdefault('choice_info', {})['per_step_timing_stats'] =\
per_step_timing_stats

self.history.append(working_output)
return working_output['chosen_action'], working_output

def reset_history(self) -> None:
self.history.clear()
for step in self.steps:
if hasattr(step, 'reset_history'):
step.reset_history()

def update_history(self, chosen_action=None, **annotations) -> None:
"""Annotate the most recent history entry with the action as it
was actually executed in the environment (which may differ from
the chosen action, e.g. truncated or partially failed plans).
Extra keyword arguments are merged into the entry as additional
annotations (e.g. `failed_actions`)."""
if self.history:
if chosen_action is not None:
self.history[-1]['executed_action'] = chosen_action
self.history[-1].update(annotations)
for step in self.steps:
if hasattr(step, 'update_history'):
step.update_history(chosen_action)
Loading