-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmain.py
More file actions
101 lines (90 loc) · 4.28 KB
/
Copy pathmain.py
File metadata and controls
101 lines (90 loc) · 4.28 KB
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
import os
import sys
import torch
import random
import argparse
import numpy as np
from dataloader import get_dataloader
from model import Trainer
def parse_args():
parser = argparse.ArgumentParser()
parser.add_argument('--cuda', action='store_true')
parser.add_argument('--seed', type=int, default=2024)
parser.add_argument('--n-family', type=int, default=5, help='number of families including human and generated')
parser.add_argument('--exp-name', type=str, default='family_moe_logits', help='model experiment name')
parser.add_argument('--pred-name', type=str, default='arxiv', help='prediction file name, just for inference')
parser.add_argument('--train-path', type=str)
parser.add_argument('--val-path', type=str)
parser.add_argument('--test-path', type=str)
parser.add_argument('--pretrain-model', default='/data/shiyuhui/pretrained/roberta-base')
parser.add_argument('--batch-size', type=int, default=64)
parser.add_argument('--max-len', type=int, default=256)
parser.add_argument('--epoch', type=int, default=50)
parser.add_argument('--lr', type=float, default=1e-3)
parser.add_argument('--early-stop', type=int, default=10)
parser.add_argument('--model-save-dir', default='./params')
parser.add_argument('--test', action='store_true')
parser.add_argument('--inference',action='store_true')
parser.add_argument('--train', action='store_true')
parser.add_argument('--is-binary', action='store_true', help='True indicate binary classification,False indicate multi-classification')
parser.add_argument('--is-cl', action='store_true', help='if use contrastive learning')
parser.add_argument('--use-proxy', action='store_true', help='use the proxy module to replace white-box probability features at inference time')
parser.add_argument('--proxy-prob', type=float, default=0.0, help='probability of using proxy-generated probability features during training')
parser.add_argument('--proxy-warmup-epochs', type=int, default=0, help='number of epochs used to warm up proxy probability and MSE weight')
parser.add_argument('--use-curriculum', action='store_true', help='linearly increase proxy probability and MSE weight during warmup')
parser.add_argument('--mse-weight', type=float, default=0.0, help='weight for MSE loss between proxy features and white-box probability features')
return parser.parse_args()
def set_seed(seed):
random.seed(seed)
os.environ['PYTHONHASHSEED'] = str(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed(seed)
torch.backends.cudnn.benchmark = False
torch.backends.cudnn.deterministic = True
def main(args):
set_seed(args.seed)
device = 'cuda' if args.cuda and torch.cuda.is_available() else 'cpu'
if not os.path.isdir(args.model_save_dir):
os.makedirs(args.model_save_dir)
model_save_path = os.path.join(args.model_save_dir, f'params_{args.exp_name}.pt')
label2id = {
'human': 0,
'generated': 1,
'llama': 1,
'mistral': 2,
'gemma': 3,
'Qwen2.5': 4,
}
train_dataloader = get_dataloader(args.train_path, args.pretrain_model, args.batch_size, args.max_len, label2id, shuffle=True) if not args.test else None
val_dataloader = get_dataloader(args.val_path, args.pretrain_model, args.batch_size, args.max_len, label2id, shuffle=False) if not args.test else None
test_dataloader = get_dataloader(args.test_path, args.pretrain_model, args.batch_size, args.max_len, label2id, shuffle=False)
trainer = Trainer(
device,
args.pretrain_model,
train_dataloader,
val_dataloader,
test_dataloader,
args.epoch,
args.lr,
args.early_stop,
model_save_path,
args.n_family,
args.is_cl,
args.is_binary,
proxy_prob=args.proxy_prob,
proxy_warmup_epochs=args.proxy_warmup_epochs,
mse_weight=args.mse_weight,
use_proxy_inference=args.use_proxy,
use_curriculum=args.use_curriculum,
)
if not args.test:
trainer.train()
else:
trainer.model.load_state_dict(torch.load(model_save_path))
results = trainer.test(test_dataloader)
print(results)
return 0
if __name__ == '__main__':
args = parse_args()
sys.exit(main(args))