MCPcopy Create free account
hub / github.com/OpenMeshLab/MeshXL / generate

Method generate

models/x_mesh_xl/get_model.py:141–165  ·  view source on GitHub ↗
(self, data_dict: dict=None, num_return_sequences: int=8, generation_config: dict=dict())

Source from the content-addressed store, hash-verified

139
140 @torch.no_grad()
141 def generate(self, data_dict: dict=None, num_return_sequences: int=8, generation_config: dict=dict()) -> dict:
142
143 net_device = next(self.parameters()).device
144 max_length = 8191
145 output_ids = torch.ones(num_return_sequences, max_length).long().to(net_device) * self.eos_token_id
146
147 # batch x ntokens
148 results = self.transformer.generate(
149 inputs_embeds=data_dict['prefix_embeds'],
150 max_length=max_length-1,
151 num_return_sequences=num_return_sequences,
152 bos_token_id=self.bos_token_id,
153 eos_token_id=self.eos_token_id,
154 pad_token_id=self.eos_token_id,
155 **generation_config
156 )
157 output_ids[:, :results.shape[1]] = results
158
159 # discard <bos> and <eos> tokens to pad tokens
160 output_ids = output_ids[:, :-1]
161 output_ids[output_ids == self.eos_token_id] = self.tokenizer.pad_id
162
163 decoder_output = self.tokenizer.detokenize(input_ids=output_ids)
164
165 return decoder_output
166
167
168

Callers 1

forwardMethod · 0.95

Calls 1

detokenizeMethod · 0.45

Tested by

no test coverage detected