| 59 | |
| 60 | |
| 61 | class 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 | |