MCPcopy Create free account
hub / github.com/cjrd/self-supervised-pretraining / main

Function main

utils/plot_pct_pretrain.py:122–281  ·  view source on GitHub ↗
(args)

Source from the content-addressed store, hash-verified

120
121
122def main(args):
123
124 #must pass in dataset arg for proper results
125
126 os.makedirs(args.out_dir, exist_ok=True)
127 # setup plots
128 sns.set_style('darkgrid')
129 sns.set()
130
131 frames = [] #array that collects dataframes from each file
132 if(args.dataset == "all"):
133 dataset_type = "*"
134 else:
135 dataset_type = args.dataset
136
137 #gets all files that start with "*pct" for the corresponding dataset
138 result_files = glob.glob(os.path.join(args.results_dir, dataset_type + "*pct*.json"), recursive=True)
139 result_files.append(os.path.join(args.results_dir, dataset_type + "_results.json")) #gets file with 100% pretrain data results
140 if (dataset_type != 'resisc'):
141 result_files.append(os.path.join(args.results_dir, dataset_type + "_bn_results.json"))
142 for resfile in result_files:
143 with open(resfile, 'r') as infile:
144 raw_data = json.load(infile)
145
146 print(resfile)
147
148 types = [] #used for concatenating the resultzs from each basetrained model
149 data = pd.DataFrame(raw_data.values())
150
151
152 #finds the relevant values for moco bt, and no bt
153 #mainly done to work around the baseline json file (has extra info that we don't want)
154 if (dataset_type + "_results.json") in resfile:
155 linear_data = data[data.result_type=='linear-eval']
156 linear_data = linear_data[linear_data.variant=="linear-eval-lr"]
157
158 data_moco = linear_data[data.basetrain=="moco_v2_800ep"]
159 # data_moco = data_moco[data_moco.pretrain_iters=="5000"]
160 data_moco = data_moco[data_moco.pretrain_data==dataset_type]
161 data_moco = reduce(data_moco)
162
163
164 data_nobt = linear_data[data.basetrain=="no"]
165 # data_nobt= data_nobt[data_nobt.pretrain_iters=="100000"]
166 data_nobt= data_nobt[data_nobt.pretrain_data==dataset_type]
167 data_nobt = reduce(data_nobt)
168
169
170 #appended to list
171 types.append(data_moco)
172 types.append(data_nobt)
173
174 #all concat to new dataframe
175 data = pd.concat(types, ignore_index=True)
176 print(data)
177
178 if "bn" in resfile:
179 data_bn = data[data.pretrain_iters!='0']

Callers 1

Calls 3

reduceFunction · 0.70
gen_plotsFunction · 0.70
setMethod · 0.45

Tested by

no test coverage detected