mirror of
https://github.com/wassname/Open-Assistant.git
synced 2026-09-09 11:15:08 +08:00
344: Create tasks for text labels (#381)
* Implement label task for initial prompts and replies * Resolve formatting * Include missing argument * Modify text_labels API to match new model, update DB schema accordingly * Send valid labels as part of label tasks * Send correctly formatted valid_labels list * Fix request format * Fix request details for text-frontend reply label task * Include message_id in tasks * Address review comments * Fix alembic tree
This commit is contained in:
@@ -24,6 +24,9 @@ class TaskType(str, enum.Enum):
|
||||
rank_initial_prompts = "rank_initial_prompts"
|
||||
rank_prompter_replies = "rank_prompter_replies"
|
||||
rank_assistant_replies = "rank_assistant_replies"
|
||||
label_initial_prompt = "label_initial_prompt"
|
||||
label_assistant_reply = "label_assistant_reply"
|
||||
label_prompter_reply = "label_prompter_reply"
|
||||
done = "task_done"
|
||||
|
||||
|
||||
@@ -56,6 +59,9 @@ class OasstApiClient:
|
||||
TaskType.rank_initial_prompts: protocol_schema.RankInitialPromptsTask,
|
||||
TaskType.rank_prompter_replies: protocol_schema.RankPrompterRepliesTask,
|
||||
TaskType.rank_assistant_replies: protocol_schema.RankAssistantRepliesTask,
|
||||
TaskType.label_initial_prompt: protocol_schema.LabelInitialPromptTask,
|
||||
TaskType.label_prompter_reply: protocol_schema.LabelPrompterReplyTask,
|
||||
TaskType.label_assistant_reply: protocol_schema.LabelAssistantReplyTask,
|
||||
TaskType.done: protocol_schema.TaskDone,
|
||||
}
|
||||
|
||||
|
||||
@@ -18,6 +18,9 @@ class TaskRequestType(str, enum.Enum):
|
||||
rank_initial_prompts = "rank_initial_prompts"
|
||||
rank_prompter_replies = "rank_prompter_replies"
|
||||
rank_assistant_replies = "rank_assistant_replies"
|
||||
label_initial_prompt = "label_initial_prompt"
|
||||
label_assistant_reply = "label_assistant_reply"
|
||||
label_prompter_reply = "label_prompter_reply"
|
||||
|
||||
|
||||
class User(BaseModel):
|
||||
@@ -169,6 +172,37 @@ class RankAssistantRepliesTask(RankConversationRepliesTask):
|
||||
type: Literal["rank_assistant_replies"] = "rank_assistant_replies"
|
||||
|
||||
|
||||
class LabelInitialPromptTask(Task):
|
||||
"""A task to label an initial prompt."""
|
||||
|
||||
type: Literal["label_initial_prompt"] = "label_initial_prompt"
|
||||
message_id: UUID
|
||||
prompt: str
|
||||
valid_labels: list[str]
|
||||
|
||||
|
||||
class LabelConversationReplyTask(Task):
|
||||
"""A task to label a reply to a conversation."""
|
||||
|
||||
type: Literal["label_conversation_reply"] = "label_conversation_reply"
|
||||
conversation: Conversation # the conversation so far
|
||||
message_id: UUID
|
||||
reply: str
|
||||
valid_labels: list[str]
|
||||
|
||||
|
||||
class LabelPrompterReplyTask(LabelConversationReplyTask):
|
||||
"""A task to label a prompter reply to a conversation."""
|
||||
|
||||
type: Literal["label_prompter_reply"] = "label_prompter_reply"
|
||||
|
||||
|
||||
class LabelAssistantReplyTask(LabelConversationReplyTask):
|
||||
"""A task to label an assistant reply to a conversation."""
|
||||
|
||||
type: Literal["label_assistant_reply"] = "label_assistant_reply"
|
||||
|
||||
|
||||
class TaskDone(Task):
|
||||
"""Signals to the frontend that the task is done."""
|
||||
|
||||
@@ -187,6 +221,10 @@ AnyTask = Union[
|
||||
RankConversationRepliesTask,
|
||||
RankPrompterRepliesTask,
|
||||
RankAssistantRepliesTask,
|
||||
LabelInitialPromptTask,
|
||||
LabelConversationReplyTask,
|
||||
LabelPrompterReplyTask,
|
||||
LabelAssistantReplyTask,
|
||||
]
|
||||
|
||||
|
||||
@@ -222,13 +260,6 @@ class MessageRanking(Interaction):
|
||||
ranking: conlist(item_type=int, min_items=1)
|
||||
|
||||
|
||||
AnyInteraction = Union[
|
||||
TextReplyToMessage,
|
||||
MessageRating,
|
||||
MessageRanking,
|
||||
]
|
||||
|
||||
|
||||
class TextLabel(str, enum.Enum):
|
||||
"""A label for a piece of text."""
|
||||
|
||||
@@ -256,12 +287,13 @@ class TextLabel(str, enum.Enum):
|
||||
slang = "slang"
|
||||
|
||||
|
||||
class TextLabels(BaseModel):
|
||||
class TextLabels(Interaction):
|
||||
"""A set of labels for a piece of text."""
|
||||
|
||||
type: Literal["text_labels"] = "text_labels"
|
||||
text: str
|
||||
labels: dict[TextLabel, float]
|
||||
message_id: str | None = None
|
||||
message_id: UUID
|
||||
|
||||
@property
|
||||
def has_message_id(self) -> bool:
|
||||
@@ -277,6 +309,14 @@ class TextLabels(BaseModel):
|
||||
return v
|
||||
|
||||
|
||||
AnyInteraction = Union[
|
||||
TextReplyToMessage,
|
||||
MessageRating,
|
||||
MessageRanking,
|
||||
TextLabels,
|
||||
]
|
||||
|
||||
|
||||
class SystemStats(BaseModel):
|
||||
all: int = 0
|
||||
active: int = 0
|
||||
|
||||
Reference in New Issue
Block a user