(self, args)
| 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( |
no test coverage detected