(args)
| 23 | return (output.float().cpu(), target.cpu()) if iscpu else (output.float(), target) |
| 24 | |
| 25 | def train(args): |
| 26 | print (f'training decoder with N_L={args.msg_len}') |
| 27 | dataset = MsgDataset(msg_len=args.msg_len, max_size=args.max_size) |
| 28 | trainloader = DataLoader(dataset, batch_size=min(args.batch_size, dataset.__len__()), shuffle=True) |
| 29 | testloader = DataLoader(dataset, batch_size=min(args.batch_size, dataset.__len__()), shuffle=False) |
| 30 | model = CLIPWatermarker(msg_len=args.msg_len).to(DEVICE) |
| 31 | optimizer = optim.Adam(model.msg_decoder.parameters(), lr=args.lr) |
| 32 | loss_fn = nn.BCEWithLogitsLoss() |
| 33 | all_epochs = trange(1, args.num_epochs + 1) |
| 34 | |
| 35 | for eIdx in all_epochs: |
| 36 | loss_value = 0 |
| 37 | for batch in trainloader: |
| 38 | optimizer.zero_grad() |
| 39 | loss = loss_fn(*single_step_inference(batch, model)) |
| 40 | loss.backward() |
| 41 | optimizer.step() |
| 42 | loss_value += loss.item() |
| 43 | |
| 44 | score = test(args, model, testloader) |
| 45 | model.train() |
| 46 | |
| 47 | all_epochs.set_description(f'BAcc : {score} | Loss : {loss_value / dataset.__len__()}') |
| 48 | |
| 49 | if args.save: |
| 50 | os.makedirs(args.sdir, exist_ok=True) |
| 51 | sdict = { |
| 52 | 'model' : model.eval().cpu().state_dict(), |
| 53 | **{key : getattr(dataset, key) for key in ['data', 'b2t_maps']} |
| 54 | } |
| 55 | torch.save(sdict, os.path.join(args.sdir, f'CLIP-MsgDecoder-{args.msg_len}.pkl')) |
| 56 | |
| 57 | def test(args, model=None, testloader=None): |
| 58 | if model is None: |
nothing calls this directly
no test coverage detected