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

Method generate

models/mesh_xl/get_model.py:170–193  ·  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

168
169 @torch.no_grad()
170 def generate(self, data_dict: dict=None, num_return_sequences: int=8, generation_config: dict=dict()) -> dict:
171
172 net_device = next(self.parameters()).device
173 max_length = 8192
174 output_ids = torch.ones(num_return_sequences, max_length).long().to(net_device) * self.eos_token_id
175
176 # batch x ntokens
177 results = self.transformer.generate(
178 max_new_tokens=max_length-1,
179 num_return_sequences=num_return_sequences,
180 bos_token_id=self.bos_token_id,
181 eos_token_id=self.eos_token_id,
182 pad_token_id=self.eos_token_id,
183 **generation_config
184 )
185 output_ids[:, :results.shape[1]] = results
186
187 # discard <bos> and <eos> tokens to pad tokens
188 output_ids = output_ids[:, 1: -1]
189 output_ids[output_ids == self.eos_token_id] = self.tokenizer.pad_id
190
191 decoder_output = self.tokenizer.detokenize(input_ids=output_ids)
192
193 return decoder_output
194
195
196

Callers 1

forwardMethod · 0.95

Calls 1

detokenizeMethod · 0.45

Tested by

no test coverage detected