MCPcopy Create free account
hub / github.com/UIC-Liu-Lab/CPT / prepare_sequence_posttrain

Function prepare_sequence_posttrain

utils/utils.py:69–119  ·  view source on GitHub ↗
(args)

Source from the content-addressed store, hash-verified

67
68
69def prepare_sequence_posttrain(args):
70 with open(os.path.join('sequences', args.sequence_file), 'r') as f:
71 datas = f.readlines()[args.idrandom]
72 data = datas.split()
73
74 args.task_name = data
75 args.current_dataset_name = data[args.pt_task]
76
77 if "cpt_datasets" in args.sequence_file:
78 args.dataset_name = 'pt'
79 else:
80 raise NotImplementedError(
81 f"The current sequence file {args.sequence_file} is not supported yet!")
82
83 output = args.base_dir + "/seq" + str(args.idrandom) + "/seed" + str(args.seed) + "/" + str(
84 args.baseline) + '/' + str(args.dataset_name) + '/' + str(data[args.pt_task]) + "_roberta/"
85 ckpt = args.base_dir + "/seq" + str(args.idrandom) + "/seed" + str(args.seed) + "/" + str(
86 args.baseline) + '/' + str(args.dataset_name) + '/' + str(data[args.pt_task - 1]) + "_roberta/"
87
88 if args.pt_task > 0:
89 args.prev_output = args.base_dir + "/seq" + str(args.idrandom) + "/seed" + str(args.seed) + "/" + str(
90 args.baseline) + '/' + str(args.dataset_name) + '/' + str(data[args.pt_task - 1]) + "_roberta/"
91 else:
92 args.prev_output = ''
93 args.task = args.pt_task
94
95 args.output_dir = output
96
97 args.saved_output_dir = [args.base_dir + "/seq" + str(args.idrandom) + "/seed" + str(args.seed) + "/" + str(
98 args.baseline) + '/' + str(args.dataset_name) + '/' + str(data[t]) + "_roberta/" for t in
99 range(args.pt_task + 1)]
100
101 if args.task == 0: # no pre-trained for the first
102 args.model_name_or_path = "roberta-base"
103 args.adapter_path = "None"
104 else:
105 args.model_name_or_path = ckpt
106 args.adapter_path = ckpt + str(args.seed) + ".model"
107
108 if 'cpt' in args.baseline:
109 args.model_name_or_path = 'roberta-base'
110
111 print('saved_output_dir: ', args.saved_output_dir)
112 print('output_dir: ', args.output_dir)
113 print('prev_output: ', args.prev_output)
114 print('dataset_name: ', args.dataset_name)
115 print('current_dataset_name: ', args.current_dataset_name)
116 print('model_name_or_path: ', args.model_name_or_path)
117 print('adapter_path: ', args.adapter_path)
118
119 return args
120
121
122def lookfor_model_finetune(args):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected