Skip to content

Commit 2ffed37

Browse files
authored
Merge pull request #394 from UKP-SQuARE/skills
Pass all kwargs to models
2 parents f21ff42 + c20e145 commit 2ffed37

18 files changed

Lines changed: 130 additions & 88 deletions

File tree

evaluator/requirements.txt

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,5 @@
11
square-elk-json-formatter==0.0.3
2-
square-skill-api==0.0.35
2+
square-skill-api==0.0.37
33
square-auth==0.0.14
44
uvicorn>=0.15.0
55
fastapi>=0.70.0

skill-manager/requirements.txt

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,5 @@
11
square-elk-json-formatter==0.0.3
2-
square-skill-api==0.0.35
2+
square-skill-api==0.0.37
33
square-auth==0.0.14
44
uvicorn>=0.15.0
55
fastapi>=0.70.0

skill-manager/skill_manager/models.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -65,7 +65,7 @@ class Skill(MongoModel):
6565
description="A description of the skill, for example describing its pipeline.",
6666
)
6767
default_skill_args: Optional[Dict] = Field(
68-
None,
68+
{},
6969
description="A dictionary holding key-value pairs that should always be sent to the skill as input. This allows to use the same skill implementataion in different ways.",
7070
)
7171
published: bool = Field(

skill-manager/skill_manager/routers/skill.py

Lines changed: 22 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -6,7 +6,6 @@
66
from threading import Thread
77
from typing import Dict, List
88

9-
import requests_cache
109
from bson import ObjectId
1110
from fastapi import APIRouter, Depends, Request
1211
from square_auth.auth import Auth
@@ -25,6 +24,7 @@
2524
from skill_manager.keycloak_api import KeycloakAPI
2625
from skill_manager.models import Prediction, Skill, SkillType
2726
from skill_manager.routers import client_credentials
27+
from skill_manager.utils import merge_dicts
2828

2929
logger = logging.getLogger(__name__)
3030

@@ -236,13 +236,29 @@ async def query_skill(
236236
query = query_request.query
237237
user_id = query_request.user_id
238238

239-
skill = await get_skill_if_authorized(request, skill_id=id, write_access=False)
239+
skill: Skill = await get_skill_if_authorized(
240+
request, skill_id=id, write_access=False
241+
)
240242
query_request.skill = json.loads(skill.json())
241243

242-
default_skill_args = skill.default_skill_args
243-
if default_skill_args is not None:
244-
# add default skill args, potentially overwrite with query.skill_args
245-
query_request.skill_args = {**default_skill_args, **query_request.skill_args}
244+
# merge kargs with kwargs in request
245+
for kwargs_key in [
246+
"model_kwargs",
247+
"task_kwargs",
248+
"preprocessing_kwargs",
249+
"explain_kwargs",
250+
"attack_kwargs",
251+
]:
252+
# overwrite kwargs from query_request with default kwargs
253+
kwargs = merge_dicts(
254+
skill.default_skill_args.pop(kwargs_key, {}),
255+
getattr(query_request, kwargs_key),
256+
)
257+
# set kwargs in query_request
258+
setattr(query_request, kwargs_key, kwargs)
259+
query_request.skill_args = merge_dicts(
260+
skill.default_skill_args, query_request.skill_args
261+
)
246262

247263
headers = {"Authorization": f"Bearer {token}"}
248264
if request.headers.get("Cache-Control"):
Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,6 @@
1+
def merge_dicts(*dicts):
2+
"""Merge multiple dictionaries into one. Overwrites values from left to right."""
3+
merged = {}
4+
for d in dicts:
5+
merged.update(d)
6+
return merged

skill-manager/tests/conftest.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -92,7 +92,7 @@ def skill_init(
9292
user_id=user_id,
9393
description=description,
9494
published=published,
95-
default_skill_args=default_skill_args,
95+
default_skill_args={} if default_skill_args is None else default_skill_args,
9696
**kwargs,
9797
)
9898
if not skill.id:

skill-manager/tests/test_api.py

Lines changed: 16 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -503,7 +503,11 @@ def test_query_skill_with_default_skill_args(
503503
):
504504
test_realm = "test-realm"
505505
test_user = "test-user"
506-
default_skill_args = {"adapter": "my-adapter", "context": "default context"}
506+
default_skill_args = {
507+
"adapter": "my-adapter",
508+
"context": "default context",
509+
"model_kwargs": {"model_foo": "model_bar"},
510+
}
507511
response, skill = create_skill_via_api(
508512
pers_client,
509513
token_factory,
@@ -540,10 +544,19 @@ def test_query_skill_with_default_skill_args(
540544
headers=dict(Authorization="Bearer " + token),
541545
)
542546

543-
actual_skill_query_body = json.loads(responses.calls[0].request.body)["skill_args"]
547+
actual_request_body = json.loads(responses.calls[0].request.body)
548+
549+
# model_kwargs is supposed to be removed from the skill_args and parsed separately
550+
TestCase().assertDictEqual(
551+
actual_request_body["model_kwargs"], default_skill_args.pop("model_kwargs")
552+
)
553+
554+
# remaining args should end up in skill_args
544555
expected_skill_query_body = default_skill_args
545556
expected_skill_query_body["context"] = query_context["context"]
546-
TestCase().assertDictEqual(actual_skill_query_body, expected_skill_query_body)
557+
TestCase().assertDictEqual(
558+
actual_request_body["skill_args"], expected_skill_query_body
559+
)
547560

548561

549562
@responses.activate

skill-manager/tests/test_utils.py

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,15 @@
1+
from skill_manager.utils import merge_dicts
2+
3+
4+
def test_merge_dicts():
5+
d1 = {"a": 1}
6+
d2 = {"b": 2}
7+
merged_dicts = merge_dicts(d1, d2)
8+
assert merged_dicts == {"a": 1, "b": 2}
9+
10+
11+
def test_overwrite_merge_dicts():
12+
d1 = {"a": 1}
13+
d2 = {"a": 2}
14+
merged_dicts = merge_dicts(d1, d2)
15+
assert merged_dicts == {"a": 2}

skills/Dockerfile

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,7 @@ COPY requirements.txt ./
1212
RUN pip install -r requirements.txt
1313

1414
COPY main.py main.py
15+
COPY utils.py utils.py
1516
ARG skill
1617
COPY ./$skill/skill.py skill.py
1718
COPY logging.conf logging.conf

skills/boolq/skill.py

Lines changed: 5 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,8 @@
33
from square_model_client import SQuAREModelClient
44
from square_skill_api.models import QueryOutput, QueryRequest
55

6+
from utils import extract_model_kwargs_from_request
7+
68
logger = logging.getLogger(__name__)
79

810
square_model_client = SQuAREModelClient()
@@ -12,18 +14,15 @@ async def predict(request: QueryRequest) -> QueryOutput:
1214
"""Predicts yes/no for a boolean question with context"""
1315
query = request.query
1416
context = request.skill_args["context"]
15-
explain_kwargs = request.explain_kwargs or {}
16-
attack_kwargs = request.attack_kwargs or {}
17+
18+
model_request_kwargs = extract_model_kwargs_from_request(request)
1719

1820
prepared_input = [[context, query]]
1921

2022
model_request = {
2123
"input": prepared_input,
22-
"preprocessing_kwargs": {},
23-
"model_kwargs": {},
2424
"adapter_name": request.skill_args["adapter"],
25-
"explain_kwargs": explain_kwargs,
26-
"attack_kwargs": attack_kwargs,
25+
**model_request_kwargs,
2726
}
2827
model_response = await square_model_client(
2928
model_name=request.skill_args["base_model"],

0 commit comments

Comments
 (0)