MCPcopy Create free account
hub / github.com/agentscope-ai/Trinity-RFT / generate

Method generate

trinity/common/models/tinker_model.py:51–121  ·  view source on GitHub ↗

Generate a responses from a prompt in async.

(self, prompt: str, **kwargs)

Source from the content-addressed store, hash-verified

49 )
50
51 async def generate(self, prompt: str, **kwargs) -> Sequence[Experience]:
52 """Generate a responses from a prompt in async."""
53 if self.tokenizer is None:
54 await self._initialize_tokenizer()
55
56 returned_seq, is_valid = self._handle_prompt_truncation(prompt, **kwargs)
57 if not is_valid:
58 return returned_seq # is_valid is False: returned_seq is a list of dummy experiences
59 token_ids = returned_seq # is_valid is True: returned_seq is prompt's token_ids
60
61 with_chat_completion = kwargs.get("with_chat_completion", False)
62 if with_chat_completion:
63 create_time = int(time.time())
64 output = await self._generate_internal(prompt={"prompt_token_ids": token_ids}, **kwargs)
65 logprobs = kwargs.get("logprobs", self.config.logprobs)
66 return_logprobs = logprobs is not None and logprobs is not False
67 experiences = [
68 Experience(
69 tokens=torch.tensor(token_ids + sequence.tokens, dtype=torch.int32),
70 logprobs=(
71 torch.tensor(sequence.logprobs, dtype=torch.float32)
72 if return_logprobs
73 else torch.tensor([], dtype=torch.float32)
74 ),
75 prompt_length=len(token_ids),
76 prompt_text=self.tokenizer.decode(token_ids),
77 response_text=self.tokenizer.decode(sequence.tokens),
78 )
79 for sequence in output.sequences
80 ]
81 if with_chat_completion:
82 from openai.types.chat.chat_completion import (
83 ChatCompletion,
84 ChatCompletionMessage,
85 ChatCompletionTokenLogprob,
86 Choice,
87 ChoiceLogprobs,
88 )
89
90 return_token_ids = kwargs.get("return_token_ids", False)
91 chat_completion = ChatCompletion(
92 id="",
93 choices=[
94 Choice(
95 finish_reason=sequence.stop_reason,
96 index=i,
97 logprobs=ChoiceLogprobs(
98 content=[
99 ChatCompletionTokenLogprob(
100 token=self.tokenizer.decode(token_id),
101 logprob=logprob,
102 top_logprobs=[],
103 )
104 for token_id, logprob in zip(sequence.tokens, sequence.logprobs)
105 ]
106 ),
107 message=ChatCompletionMessage(
108 content=self.tokenizer.decode(sequence.tokens), role="assistant"

Callers 1

chatMethod · 0.95

Calls 6

_initialize_tokenizerMethod · 0.95
_generate_internalMethod · 0.95
ExperienceClass · 0.90
getMethod · 0.45
decodeMethod · 0.45

Tested by

no test coverage detected