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

Class MeshXL

models/mesh_xl/get_model.py:9–193  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

7
8
9class MeshXL(nn.Module):
10
11 def train(self, mode: bool = True):
12 super().train(mode)
13 return self
14
15 def __init__(self, args):
16 super().__init__()
17
18 self.tokenizer = MeshTokenizer(args)
19
20 # causal LM model initialization
21 self.vocab_size = self.tokenizer.codebook_size + 3
22 self.bos_token_id = self.tokenizer.codebook_size
23 self.eos_token_id = self.tokenizer.codebook_size + 1
24 self.pad_token_id = self.tokenizer.codebook_size + 2
25
26 config = AutoConfig.from_pretrained(
27 args.llm,
28 n_positions=8192,
29 max_position_embeddings=8192,
30 vocab_size=self.vocab_size,
31 bos_token_id=self.bos_token_id,
32 eos_token_id=self.eos_token_id,
33 pad_token_id=self.pad_token_id
34 )
35
36 config.word_embed_proj_dim = config.hidden_size
37 self.transformer = AutoModelForCausalLM.from_pretrained(
38 args.llm,
39 config=config,
40 ignore_mismatched_sizes=True
41 )
42 self.transformer.to_bettertransformer()
43
44 # setting status for all parameters
45 self.train()
46
47
48 def forward(
49 self,
50 data_dict: dict=None,
51 is_eval: bool=False,
52 is_generate: bool=False,
53 num_return_sequences: int=8,
54 generation_config: Dict=dict(
55 do_sample=True,
56 top_k=50,
57 top_p=0.95,
58 # no_repeat_ngram_size=9,
59 )
60 ) -> dict:
61
62 if not is_eval:
63 return self.train_one_step(data_dict)
64
65 if is_eval and not is_generate:
66 return self.perplexity(data_dict)

Callers 1

get_modelFunction · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected