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

Function main

utils/plot_pct_pretrain_modified.py:147–306  ·  view source on GitHub ↗
(args)

Source from the content-addressed store, hash-verified

145
146
147def main(args):
148
149 #must pass in dataset arg for proper results
150
151 os.makedirs(args.out_dir, exist_ok=True)
152 # setup plots
153 sns.set_style('darkgrid')
154 sns.set()
155
156 frames = [] #array that collects dataframes from each file
157 if(args.dataset == "all"):
158 dataset_type = "*"
159 else:
160 dataset_type = args.dataset
161
162 #gets all files that start with "*pct" for the corresponding dataset
163 result_files = glob.glob(os.path.join(args.results_dir, dataset_type + "*pct*.json"), recursive=True)
164 result_files.append(os.path.join(args.results_dir, dataset_type + "_results.json")) #gets file with 100% pretrain data results
165 if (dataset_type != 'resisc'):
166 result_files.append(os.path.join(args.results_dir, dataset_type + "_bn_results.json"))
167 for resfile in result_files:
168 with open(resfile, 'r') as infile:
169 raw_data = json.load(infile)
170
171 print(resfile)
172
173 types = [] #used for concatenating the resultzs from each basetrained model
174 data = pd.DataFrame(raw_data.values())
175
176
177 #finds the relevant values for moco bt, and no bt
178 #mainly done to work around the baseline json file (has extra info that we don't want)
179 if (dataset_type + "_results.json") in resfile:
180 linear_data = data[data.result_type=='linear-eval']
181 linear_data = linear_data[linear_data.variant=="linear-eval-lr"]
182
183 data_moco = linear_data[data.basetrain=="moco_v2_800ep"]
184 # data_moco = data_moco[data_moco.pretrain_iters=="5000"]
185 data_moco = data_moco[data_moco.pretrain_data==dataset_type]
186 data_moco = reduce(data_moco)
187
188
189 data_nobt = linear_data[data.basetrain=="no"]
190 # data_nobt= data_nobt[data_nobt.pretrain_iters=="100000"]
191 data_nobt= data_nobt[data_nobt.pretrain_data==dataset_type]
192 data_nobt = reduce(data_nobt)
193
194
195 #appended to list
196 types.append(data_moco)
197 types.append(data_nobt)
198
199 #all concat to new dataframe
200 data = pd.concat(types, ignore_index=True)
201 print(data)
202
203 if "bn" in resfile:
204 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