MCPcopy Create free account
hub / github.com/NarcissusEx/GuardSplat / train

Function train

make_decoder.py:25–55  ·  view source on GitHub ↗
(args)

Source from the content-addressed store, hash-verified

23 return (output.float().cpu(), target.cpu()) if iscpu else (output.float(), target)
24
25def 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
57def test(args, model=None, testloader=None):
58 if model is None:

Callers

nothing calls this directly

Calls 5

__len__Method · 0.95
MsgDatasetClass · 0.90
CLIPWatermarkerClass · 0.90
single_step_inferenceFunction · 0.85
testFunction · 0.85

Tested by

no test coverage detected