Skip to content
Merged
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
14 changes: 14 additions & 0 deletions src/unitxt/formats.py
Original file line number Diff line number Diff line change
Expand Up @@ -528,6 +528,9 @@ class HFSystemFormat(ChatAPIFormat):

See more details in https://huggingface.co/docs/transformers/main/en/chat_templating

If the tokenizer of the model does not define a chat template, a Jinja chat template can be passed explicitly:
``HFSystemFormat(model_name=..., chat_kwargs_dict={"chat_template": "<template>"})``

"""

model_name: str
Expand All @@ -540,6 +543,17 @@ def prepare(self):
from transformers import AutoTokenizer

self.tokenizer = AutoTokenizer.from_pretrained(self.model_name)
if (
getattr(self.tokenizer, "chat_template", None) is None
and "chat_template" not in self.chat_kwargs_dict
):
raise UnitxtError(
f"HFSystemFormat cannot be used with model '{self.model_name}' because its tokenizer "
"does not define a chat template (no 'chat_template' in its tokenizer_config.json). "
"Either use a model whose tokenizer has a chat template, pass a Jinja chat template "
"explicitly with chat_kwargs_dict={'chat_template': '<template>'}, "
"or use a format that does not depend on the tokenizer, such as SystemFormat."
)

def _format_instance_to_source(
self,
Expand Down
44 changes: 44 additions & 0 deletions tests/library/test_formats.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
from unitxt.api import load_dataset
from unitxt.card import TaskCard
from unitxt.collections_operators import Wrap
from unitxt.error_utils import UnitxtError
from unitxt.formats import (
ChatAPIFormat,
GraniteDocumentsFormat,
Expand Down Expand Up @@ -338,6 +339,49 @@ def test_hf_system_format(self):
tester=self,
)

def test_hf_system_format_without_chat_template(self):
# The gpt2 tokenizer does not define a chat template
with self.assertRaises(UnitxtError) as cm:
HFSystemFormat(model_name="openai-community/gpt2")
self.assertIn("does not define a chat template", str(cm.exception))
self.assertIn("chat_kwargs_dict", str(cm.exception))

# A chat template can be passed explicitly instead
system_format = HFSystemFormat(
model_name="openai-community/gpt2",
chat_kwargs_dict={
"chat_template": "{% for message in messages %}<|{{ message['role'] }}|>\n{{ message['content'] }}\n{% endfor %}"
"{% if add_generation_prompt %}<|assistant|>\n{% endif %}"
},
)

inputs = [
{
"source": "1+1",
"target": "2",
"instruction": "solve the math exercises",
"demos": [],
"input_fields": {},
"target_prefix": "The answer is ",
"system_prompt": "You are a smart assistant.",
},
]
targets = [
{
"target": "2",
"input_fields": {},
"source": "<|system|>\nYou are a smart assistant.\nsolve the math exercises\n<|user|>\n1+1\n<|assistant|>\nThe answer is ",
"demos": [],
},
]

check_operator(
operator=system_format,
inputs=inputs,
targets=targets,
tester=self,
)

def test_granite_documents_format(self):
inputs = [
{
Expand Down
24 changes: 24 additions & 0 deletions tests/library/test_metrics.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,6 @@
import random
import unittest
from importlib.util import find_spec
from math import isnan
from typing import Dict, List

Expand Down Expand Up @@ -89,6 +91,13 @@

logger = get_logger()

# ReflectionToolCallingMetric and ReflectionToolCallingMetricSyntactic require llmevalkit,
# an internal package that CI installs only when repository secrets are available
# (not for pull requests from forks or from Dependabot).
requires_llmevalkit = unittest.skipUnless(
find_spec("llmevalkit") is not None, "requires llmevalkit, which is not installed"
)

# values of inputs that are common to grouped_mean type InstanceMetric
GROUPED_INSTANCE_PREDICTIONS = [
"A B",
Expand Down Expand Up @@ -1654,6 +1663,7 @@ def test_tool_calling_metric(self):
outputs[0]["score"]["global"]["argument_schema_validation"], 0.0
)

@requires_llmevalkit
def test_reflection_tool_calling_metric(self):
unitxt.settings.mock_inference_mode = True
metric = ReflectionToolCallingMetric()
Expand Down Expand Up @@ -1730,6 +1740,7 @@ def test_reflection_tool_calling_metric(self):
]
)

@requires_llmevalkit
def test_partial_value_precision_enum_violations_real_static_only(self):
"""Test partial value precision when some parameters have invalid enum values."""
metric = ReflectionToolCallingMetricSyntactic()
Expand Down Expand Up @@ -1790,6 +1801,7 @@ def test_partial_value_precision_enum_violations_real_static_only(self):
result["metrics"]["missing_required_parameter"]["valid"], True
)

@requires_llmevalkit
def test_reflection_tool_calling_metric_reduce(self):
# Instance 1: valid call
instance1 = {
Expand Down Expand Up @@ -1942,6 +1954,7 @@ def test_reflection_tool_calling_metric_reduce(self):
reduced["semantic_agentic_constraints_satisfaction"], acs_expected
)

@requires_llmevalkit
def test_reflection_tool_calling_metric_syntactic_reduce(self):
from unitxt.metrics import ReflectionToolCallingMetricSyntactic

Expand Down Expand Up @@ -2031,6 +2044,7 @@ def mean_valid(name: str) -> float:
reduced_shuffled = metric.reduce(shuffled)
self.assertEqual(reduced, reduced_shuffled)

@requires_llmevalkit
def test_tool_calling_metric_syntactic_reflector(self):
metric = ReflectionToolCallingMetricSyntactic()
tools_data = {
Expand Down Expand Up @@ -2201,6 +2215,7 @@ def test_tool_calling_metric_syntactic_reflector(self):
# schema validation can still pass even if there are type errors
self.assertTrue(outputs["metrics"]["json_schema_violation"]["valid"])

@requires_llmevalkit
def test_overall_valid_success_real_map(self):
metric = ReflectionToolCallingMetricSyntactic()
# Create sample inputs
Expand Down Expand Up @@ -2248,6 +2263,7 @@ def test_overall_valid_success_real_map(self):
self.assertTrue(result["metrics"]["non_existent_function"]["valid"])
self.assertTrue(result["metrics"]["missing_required_parameter"]["valid"])

@requires_llmevalkit
def test_non_existent_function_real_map(self):
metric = ReflectionToolCallingMetricSyntactic()
# Create sample inputs with wrong function name
Expand Down Expand Up @@ -2282,6 +2298,7 @@ def test_non_existent_function_real_map(self):
self.assertFalse(result["metrics"]["non_existent_function"]["valid"])
self.assertTrue(result["metrics"]["missing_required_parameter"]["valid"])

@requires_llmevalkit
def test_missing_required_parameter_real_map(self):
metric = ReflectionToolCallingMetricSyntactic()
# Create sample inputs with missing required parameter
Expand Down Expand Up @@ -2322,6 +2339,7 @@ def test_missing_required_parameter_real_map(self):
self.assertFalse(result["metrics"]["missing_required_parameter"]["valid"])
self.assertTrue(result["metrics"]["allowed_values_violation"]["valid"])

@requires_llmevalkit
def test_non_existent_parameter_real_map(self):
metric = ReflectionToolCallingMetricSyntactic()
# Create sample inputs with extra undefined parameter
Expand Down Expand Up @@ -2355,6 +2373,7 @@ def test_non_existent_parameter_real_map(self):
self.assertFalse(result["overall_valid"], False)
self.assertFalse(result["metrics"]["non_existent_parameter"]["valid"])

@requires_llmevalkit
def test_allowed_values_violation(self):
metric = ReflectionToolCallingMetricSyntactic()
# Create sample inputs with invalid enum value
Expand Down Expand Up @@ -2401,6 +2420,7 @@ def test_allowed_values_violation(self):
self.assertFalse(result["metrics"]["allowed_values_violation"]["valid"])
self.assertTrue(result["metrics"]["incorrect_parameter_type"]["valid"])

@requires_llmevalkit
def test_json_schema_violation_specific_real_map(self):
metric = ReflectionToolCallingMetricSyntactic()
# Create sample inputs
Expand Down Expand Up @@ -2441,6 +2461,7 @@ def test_json_schema_violation_specific_real_map(self):
# json_schema_violation specifically should be 1.0 because it's marked valid
self.assertTrue(result["metrics"]["json_schema_violation"]["valid"])

@requires_llmevalkit
def test_partial_recall_missing_parameters_real_map(self):
metric = ReflectionToolCallingMetricSyntactic()
"""Test partial recall score when some but not all required parameters are missing."""
Expand Down Expand Up @@ -2479,6 +2500,7 @@ def test_partial_recall_missing_parameters_real_map(self):
self.assertFalse(result["overall_valid"])
self.assertFalse(result["metrics"]["missing_required_parameter"]["valid"])

@requires_llmevalkit
def test_partial_precision_non_existent_parameters_real_map(self):
"""Test partial precision score when some parameters don't exist in the schema."""
metric = ReflectionToolCallingMetricSyntactic()
Expand Down Expand Up @@ -2524,6 +2546,7 @@ def test_partial_precision_non_existent_parameters_real_map(self):
self.assertFalse(result["metrics"]["non_existent_parameter"]["valid"])
self.assertFalse(result["overall_valid"])

@requires_llmevalkit
def test_partial_value_precision_type_errors_real_map(self):
"""Test partial value precision when some parameters have incorrect types."""
metric = ReflectionToolCallingMetricSyntactic()
Expand Down Expand Up @@ -2572,6 +2595,7 @@ def test_partial_value_precision_type_errors_real_map(self):
result["metrics"]["incorrect_parameter_type"]["valid"], False
)

@requires_llmevalkit
def test_partial_value_precision_enum_violations_real_map(self):
"""Test partial value precision when some parameters have invalid enum values."""
metric = ReflectionToolCallingMetricSyntactic()
Expand Down
Loading