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

Function test

make_decoder.py:57–76  ·  view source on GitHub ↗
(args, model=None, testloader=None)

Source from the content-addressed store, hash-verified

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:
59 print (f'testing decoder with N_L={args.msg_len}')
60 sdict = torch.load(os.path.join(args.sdir, f'CLIP-MsgDecoder-{args.msg_len}.pkl'), map_location='cpu', weights_only=True)
61 model = CLIPWatermarker(msg_len=args.msg_len)
62 model.load_state_dict(sdict.pop('model'))
63 model.to(DEVICE)
64 dataset = MsgDataset(msg_len=args.msg_len, max_size=args.max_size, **sdict)
65 testloader = DataLoader(dataset, batch_size=min(args.batch_size, dataset.__len__()), shuffle=False)
66
67 with torch.no_grad():
68 model.eval()
69 outputs, targets = list(zip(*[single_step_inference(batch, model, iscpu=True) for batch in testloader]))
70 outputs, targets = torch.cat(outputs), torch.cat(targets)
71
72 if args.mode == 'train':
73 return bit_accuracy(outputs, targets)
74
75 else:
76 print (f'BAcc : {bit_accuracy(outputs, targets):.2f}')
77
78if __name__ == '__main__':
79 parser = argparse.ArgumentParser()

Callers 1

trainFunction · 0.85

Calls 5

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

Tested by

no test coverage detected