Main training program.
()
| 241 | return tokenizer |
| 242 | |
| 243 | def main(): |
| 244 | """Main training program.""" |
| 245 | |
| 246 | print('Generate Samples') |
| 247 | |
| 248 | # Disable CuDNN. |
| 249 | torch.backends.cudnn.enabled = False |
| 250 | |
| 251 | # Timer. |
| 252 | timers = Timers() |
| 253 | |
| 254 | # Arguments. |
| 255 | args = get_args() |
| 256 | |
| 257 | # Pytorch distributed. |
| 258 | initialize_distributed(args) |
| 259 | |
| 260 | # Random seeds for reproducability. |
| 261 | set_random_seed(args.seed) |
| 262 | |
| 263 | #get the tokenizer |
| 264 | tokenizer = prepare_tokenizer(args) |
| 265 | |
| 266 | # Model, optimizer, and learning rate. |
| 267 | model = setup_model(args) |
| 268 | |
| 269 | #setting default batch size to 1 |
| 270 | args.batch_size = 1 |
| 271 | |
| 272 | #generate samples |
| 273 | generate_samples(model, tokenizer, args, torch.cuda.current_device()) |
| 274 | |
| 275 | |
| 276 | if __name__ == "__main__": |
no test coverage detected