File size: 4,120 Bytes
9e14838
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
import torch
from train import Trainer
from evaluate import Detector
from utils import seed_torch


def main(args):
    #########################################
    # Phase 1: Training
    #########################################
    if args.phase == "train":
        # Initialize Trainer
        trainer = Trainer(
            mode=args.mode,
            device=args.device,
            branch=args.branch,
            train_dataset=args.train_dataset,
            label_smooth=args.label_smooth,
            ckpt=args.ckpt,
            epochs=args.epochs,
            batch_size=args.batch_size
        )
        # Start training
        trainer.train()
    #########################################
    # Phase 2: Evaluation
    #########################################
    elif args.phase == "eval":
        # Initialize Detector
        detector = Detector(
            device=args.device,
            mode=args.mode,
            train_dataset=args.train_dataset,
            pretrain=args.pretrain,
            ckpt=args.ckpt,
            batch_size=args.batch_size
        )
        # Start evaluation
        detector.evaluate_benchmark()
    ##########################################
    # Phase 3: Test on a single image
    ##########################################
    elif args.phase == "test":
        # Initialize Detector
        detector = Detector(
            device=args.device,
            mode=args.mode,
            train_dataset=args.train_dataset,
            pretrain=args.pretrain,
            ckpt=args.ckpt,
            batch_size=args.batch_size
        )
        # Test on a single image
        score = detector.scan()
    else:
        raise ValueError(f"Unknown phase: {args.phase}")


if __name__ == "__main__":
    import argparse
    parser = argparse.ArgumentParser("Co-Spy: Combining Semantic and Pixel Features to Detect Synthetic Images by AI")
    parser.add_argument("--gpu",
                        type=int,
                        default=0,
                        help="GPU id to use")
    parser.add_argument("--phase",
                        type=str,
                        default="test",
                        choices=["train", "eval", "test"],
                        help="Select the phase to run Co-Spy: train / eval / test")
    parser.add_argument("--mode",
                        type=str,
                        default="fusion",
                        choices=["branch", "fusion", "end2end"],
                        help="Select the mode of Co-Spy training")
    parser.add_argument("--train_dataset",
                        type=str,
                        default="sd-v1_4",
                        help="Training dataset")
    parser.add_argument("--branch",
                        type=str,
                        default="artifact",
                        choices=["artifact", "semantic"],
                        help="Branch detector (for branch mode)")
    parser.add_argument("--label_smooth",
                        action="store_true",
                        help="Whether to use label smoothing during training")
    parser.add_argument("--pretrain",
                        action="store_true",
                        help="Whether to use pre-trained weights for evaluation")
    parser.add_argument("--ckpt",
                        type=str,
                        default="ckpt",
                        help="Checkpoint directory")
    parser.add_argument("--epochs",
                        type=int,
                        default=20,
                        help="Number of training epochs")
    parser.add_argument("--batch_size",
                        type=int,
                        default=32,
                        help="Batch size")
    parser.add_argument("--seed",
                        type=int,
                        default=1024,
                        help="Random seed")

    args = parser.parse_args()

    # Set random seed
    seed_torch(args.seed)

    # Set GPU device
    args.device = f"cuda:{args.gpu}" if torch.cuda.is_available() else "cpu"

    # Run the experiment
    main(args)