Fetch conversation for seed data tasks, minor model fixes (#485)

* Fetch conversation for seed data, fix models, remove redundant payload type checks
This commit is contained in:
Andreas Köpf
2023-01-07 15:59:54 +01:00
committed by GitHub
parent eaefa68dea
commit 96d6717be4
9 changed files with 61 additions and 54 deletions
@@ -2,9 +2,7 @@ from fastapi import APIRouter, Depends
from oasst_backend.api import deps
from oasst_backend.api.v1 import utils
from oasst_backend.models import ApiClient
from oasst_backend.models.db_payload import MessagePayload
from oasst_backend.prompt_repository import PromptRepository
from oasst_shared.exceptions import OasstError, OasstErrorCode
from oasst_shared.schemas import protocol
from sqlmodel import Session
@@ -20,11 +18,6 @@ def get_message_by_frontend_id(
"""
pr = PromptRepository(db, api_client, user=None)
message = pr.fetch_message_by_frontend_message_id(message_id)
if not isinstance(message.payload.payload, MessagePayload):
# Unexpected message payload
raise OasstError("Invalid message", OasstErrorCode.INVALID_MESSAGE)
return utils.prepare_message(message)
-6
View File
@@ -5,9 +5,7 @@ from fastapi import APIRouter, Depends, Query
from oasst_backend.api import deps
from oasst_backend.api.v1 import utils
from oasst_backend.models import ApiClient
from oasst_backend.models.db_payload import MessagePayload
from oasst_backend.prompt_repository import PromptRepository
from oasst_shared.exceptions import OasstError, OasstErrorCode
from oasst_shared.schemas import protocol
from sqlmodel import Session
from starlette.status import HTTP_204_NO_CONTENT
@@ -55,10 +53,6 @@ def get_message(
"""
pr = PromptRepository(db, api_client, user=None)
message = pr.fetch_message(message_id)
if not isinstance(message.payload.payload, MessagePayload):
# Unexptcted message payload
raise OasstError("Invalid message", OasstErrorCode.INVALID_MESSAGE)
return utils.prepare_message(message)
+10 -14
View File
@@ -57,9 +57,7 @@ def generate_task(
logger.info("Generating a PrompterReplyTask.")
messages = pr.fetch_random_conversation("assistant")
task_messages = [
protocol_schema.ConversationMessage(
text=msg.payload.payload.text, is_assistant=(msg.role == "assistant")
)
protocol_schema.ConversationMessage(text=msg.text, is_assistant=(msg.role == "assistant"))
for msg in messages
]
@@ -70,9 +68,7 @@ def generate_task(
logger.info("Generating a AssistantReplyTask.")
messages = pr.fetch_random_conversation("prompter")
task_messages = [
protocol_schema.ConversationMessage(
text=msg.payload.payload.text, is_assistant=(msg.role == "assistant")
)
protocol_schema.ConversationMessage(text=msg.text, is_assistant=(msg.role == "assistant"))
for msg in messages
]
@@ -83,19 +79,19 @@ def generate_task(
logger.info("Generating a RankInitialPromptsTask.")
messages = pr.fetch_random_initial_prompts()
task = protocol_schema.RankInitialPromptsTask(prompts=[msg.payload.payload.text for msg in messages])
task = protocol_schema.RankInitialPromptsTask(prompts=[msg.text for msg in messages])
case protocol_schema.TaskRequestType.rank_prompter_replies:
logger.info("Generating a RankPrompterRepliesTask.")
conversation, replies = pr.fetch_multiple_random_replies(message_role="assistant")
task_messages = [
protocol_schema.ConversationMessage(
text=p.payload.payload.text,
text=p.text,
is_assistant=(p.role == "assistant"),
)
for p in conversation
]
replies = [p.payload.payload.text for p in replies]
replies = [p.text for p in replies]
task = protocol_schema.RankPrompterRepliesTask(
conversation=protocol_schema.Conversation(
messages=task_messages,
@@ -109,12 +105,12 @@ def generate_task(
task_messages = [
protocol_schema.ConversationMessage(
text=p.payload.payload.text,
text=p.text,
is_assistant=(p.role == "assistant"),
)
for p in conversation
]
replies = [p.payload.payload.text for p in replies]
replies = [p.text for p in replies]
task = protocol_schema.RankAssistantRepliesTask(
conversation=protocol_schema.Conversation(messages=task_messages),
replies=replies,
@@ -125,14 +121,14 @@ def generate_task(
message = pr.fetch_random_initial_prompts(1)[0]
task = protocol_schema.LabelInitialPromptTask(
message_id=message.id,
prompt=message.payload.payload.text,
prompt=message.text,
valid_labels=list(map(lambda x: x.value, protocol_schema.TextLabel)),
)
case protocol_schema.TaskRequestType.label_prompter_reply:
logger.info("Generating a LabelPrompterReplyTask.")
conversation, messages = pr.fetch_multiple_random_replies(max_size=1, message_role="assistant")
message = messages[0].payload.payload.text
message = messages[0].text
task = protocol_schema.LabelPrompterReplyTask(
message_id=message.id,
conversation=conversation,
@@ -143,7 +139,7 @@ def generate_task(
case protocol_schema.TaskRequestType.label_assistant_reply:
logger.info("Generating a LabelAssistantReplyTask.")
conversation, messages = pr.fetch_multiple_random_replies(max_size=1, message_role="prompter")
message = messages[0].payload.payload.text
message = messages[0].text
task = protocol_schema.LabelAssistantReplyTask(
message_id=message.id,
conversation=conversation,
+2 -11
View File
@@ -1,19 +1,14 @@
from http import HTTPStatus
from uuid import UUID
from oasst_backend.models import Message
from oasst_backend.models.db_payload import MessagePayload
from oasst_shared.exceptions import OasstError, OasstErrorCode
from oasst_shared.schemas import protocol
def prepare_message(m: Message) -> protocol.Message:
if not isinstance(m.payload.payload, MessagePayload):
raise OasstError("Server error", OasstErrorCode.SERVER_ERROR, HTTPStatus.INTERNAL_SERVER_ERROR)
return protocol.Message(
id=m.id,
parent_id=m.parent_id,
text=m.payload.payload.text,
text=m.text,
is_assistant=(m.role == "assistant"),
created_date=m.created_date,
)
@@ -26,10 +21,8 @@ def prepare_message_list(messages: list[Message]) -> list[protocol.Message]:
def prepare_conversation(messages: list[Message]) -> protocol.Conversation:
conv_messages = []
for message in messages:
if not isinstance(message.payload.payload, MessagePayload):
raise OasstError("Server error", OasstErrorCode.SERVER_ERROR, HTTPStatus.INTERNAL_SERVER_ERROR)
conv_messages.append(
protocol.ConversationMessage(text=message.payload.payload.text, is_assistant=(message.role == "assistant"))
protocol.ConversationMessage(text=message.text, is_assistant=(message.role == "assistant"))
)
return protocol.Conversation(messages=conv_messages)
@@ -38,8 +31,6 @@ def prepare_conversation(messages: list[Message]) -> protocol.Conversation:
def prepare_tree(tree: list[Message], tree_id: UUID) -> protocol.MessageTree:
tree_messages = []
for message in tree:
if not isinstance(message.payload.payload, MessagePayload):
raise OasstError("Server error", OasstErrorCode.SERVER_ERROR, HTTPStatus.INTERNAL_SERVER_ERROR)
tree_messages.append(prepare_message(message))
return protocol.MessageTree(id=tree_id, messages=tree_messages)