(args, model=None, testloader=None)
| 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: |
| 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 | |
| 78 | if __name__ == '__main__': |
| 79 | parser = argparse.ArgumentParser() |
no test coverage detected