(num_models, config_dict, dtype)
| 19 | |
| 20 | |
| 21 | def train_shared_loss(num_models, config_dict, dtype): |
| 22 | hidden_dim = 64 |
| 23 | |
| 24 | models = [create_model(config_dict) for _ in range(num_models)] |
| 25 | data_loader = random_dataloader(model=models[0], |
| 26 | total_samples=4, |
| 27 | hidden_dim=hidden_dim, |
| 28 | device=models[0].device, |
| 29 | dtype=dtype) |
| 30 | dist.barrier() |
| 31 | for _, batch in enumerate(data_loader): |
| 32 | losses = [m.module(batch[0], batch[1]) for m in models] |
| 33 | loss = sum(l / (i + 1) for i, l in enumerate(losses)) |
| 34 | loss.backward() |
| 35 | |
| 36 | for m in models: |
| 37 | m._backward_epilogue() |
| 38 | |
| 39 | for m in models: |
| 40 | m.step() |
| 41 | |
| 42 | for m in models: |
| 43 | m.optimizer.zero_grad() |
| 44 | |
| 45 | for m in models: |
| 46 | m.destroy() |
| 47 | |
| 48 | |
| 49 | def train_independent_loss(num_models, config_dict, dtype): |
no test coverage detected