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

Class MeshXL

models/x_mesh_xl/get_model.py:61–165  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

59
60
61class MeshXL(nn.Module):
62
63 def train(self, mode: bool = True):
64 super().train(mode)
65 # self.transformer.eval()
66 # for param in self.transformer.parameters():
67 # param.requires_grad = False
68 return self
69
70 def __init__(self, args):
71 super().__init__()
72
73 self.tokenizer = MeshTokenizer(args)
74
75 # causal LM model initialization
76 self.vocab_size = self.tokenizer.codebook_size + 3
77 self.bos_token_id = self.tokenizer.codebook_size
78 self.eos_token_id = self.tokenizer.codebook_size + 1
79 self.pad_token_id = self.tokenizer.codebook_size + 2
80
81 config = AutoConfig.from_pretrained(
82 args.llm,
83 n_positions=8192,
84 max_position_embeddings=8192,
85 vocab_size=self.vocab_size,
86 bos_token_id=self.bos_token_id,
87 eos_token_id=self.eos_token_id,
88 pad_token_id=self.pad_token_id
89 )
90
91 config.word_embed_proj_dim = config.hidden_size
92 self.transformer = AutoModelForCausalLM.from_config(config=config)
93
94 try:
95 self.transformer.to_bettertransformer()
96 except:
97 pass
98
99 self.condition_encoder = ConditionEncoder(args, config.hidden_size)
100
101 # setting status for all parameters
102 self.train()
103
104
105 def forward(
106 self,
107 data_dict: dict=None,
108 is_eval: bool=False,
109 is_generate: bool=False,
110 num_return_sequences: int=8,
111 generation_config: Dict=dict(
112 do_sample=True,
113 top_k=50,
114 top_p=0.95,
115 # no_repeat_ngram_size=9,
116 )
117 ) -> dict:
118

Callers 1

get_modelFunction · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected