| 38 | |
| 39 | |
| 40 | def gen_plots(data, options): |
| 41 | |
| 42 | # linear plots |
| 43 | lin_dir = os.path.join(options['out_dir'], 'linear-eval') |
| 44 | os.makedirs(lin_dir, exist_ok=True) |
| 45 | linear_data = data[data.result_type=='linear-eval'] |
| 46 | |
| 47 | #for axes for subplots |
| 48 | i = 0 |
| 49 | j = 0 |
| 50 | |
| 51 | #creates suplots |
| 52 | #didn't use setup_plot since the aspect ratios weren't ideal for the figures |
| 53 | with sns.axes_style("ticks"): |
| 54 | fig, axes = plt.subplots(4, 4, figsize=(25, 16), sharey=False, constrained_layout=True) |
| 55 | for pretrain_data in linear_data.pretrain_data.unique(): |
| 56 | for variant in ["linear-eval-lr"]: # linear_data['variant'].unique(): |
| 57 | with sns.axes_style("ticks"): |
| 58 | |
| 59 | #checks to make sure axes indices are in range |
| 60 | if j > 3: |
| 61 | i += 1 |
| 62 | j %= 4 |
| 63 | |
| 64 | variant_data=linear_data[(linear_data.variant == variant) & (linear_data.pretrain_data == pretrain_data)] |
| 65 | # hack to "clean up" other analyses |
| 66 | variant_data = variant_data[(linear_data.basetrain == 'MoCo Random Init') | (linear_data.basetrain == 'HPT') ] |
| 67 | |
| 68 | |
| 69 | print(i,j) |
| 70 | ax1 = axes[i,j] |
| 71 | |
| 72 | #set up x scale |
| 73 | ax1.set_xscale('symlog') |
| 74 | # data = data[~((data.pretrain_iters == "5000") & (data.basetrain == "MoCo Random Init"))] |
| 75 | ax1.set_xticks(ticks=variant_data['pretrain_iters'].unique().astype(int)) |
| 76 | |
| 77 | name_map = { |
| 78 | 50: '50', |
| 79 | 500: '500', |
| 80 | 5000: '5K', |
| 81 | 50000: '50K', |
| 82 | 100000: '100K', |
| 83 | 200000: '200K', |
| 84 | 400000: '400K' |
| 85 | } |
| 86 | |
| 87 | ori_labels = variant_data['pretrain_iters'].unique().astype(int) |
| 88 | new_labels = [name_map[key] for key in ori_labels] |
| 89 | |
| 90 | ax1.set_xticklabels(labels=new_labels, minor=False, rotation=40) |
| 91 | print(variant_data['pretrain_iters'].unique()) |
| 92 | |
| 93 | #hard coded bn_data result |
| 94 | |
| 95 | #plot data |
| 96 | ax1.axhline(y=options['moco_transfers'][pretrain_data], c='black', linestyle='dashed', linewidth=3, label="MoCo Direct Transfer") |
| 97 | |