| 121 | |
| 122 | |
| 123 | def load_variation(env, args, task_num, logger): |
| 124 | variations = [] |
| 125 | if (args["set"] == "train"): |
| 126 | variations = list(env.getVariationsTrain()) |
| 127 | if task_num == 26: |
| 128 | variations = variations[:int(len(variations)/10)] |
| 129 | elif task_num == 29: |
| 130 | variations = variations[:int(len(variations)/2)] |
| 131 | elif (args["set"] == "test"): |
| 132 | variations = list(env.getVariationsTest()) |
| 133 | if True or args["cut_off"]: |
| 134 | test_len = min(10, len(variations)) |
| 135 | if task_num == 25: |
| 136 | test_len = 5 |
| 137 | elif task_num == 15: |
| 138 | test_len = 9 |
| 139 | random.seed(1) |
| 140 | random.shuffle(variations) |
| 141 | variations = variations[:test_len] |
| 142 | print(f'{task_num}: {len(variations)}') |
| 143 | elif (args["set"] == "dev"): |
| 144 | variations = list(env.getVariationsDev()) |
| 145 | variations = variations[:3] |
| 146 | elif (args["set"] == "test_mini_2"): |
| 147 | variations = list(env.getVariationsTest()) |
| 148 | # random.seed(1) |
| 149 | # random.shuffle(variations) |
| 150 | variations = variations[3:10] |
| 151 | elif (args["set"] == "test_mini"): |
| 152 | variations = list(env.getVariationsTest()) |
| 153 | # random.seed(1) |
| 154 | # random.shuffle(variations) |
| 155 | variations = variations[:3] |
| 156 | elif (args["set"] == "test_mini_mini"): |
| 157 | variations = list(env.getVariationsTest()) |
| 158 | # random.seed(1) |
| 159 | # random.shuffle(variations) |
| 160 | variations = variations[:1] |
| 161 | else: |
| 162 | logger.info("ERROR: Unknown set to evaluate on (" + str(args["set"]) + ")") |
| 163 | exit(1) |
| 164 | |
| 165 | logger.info(variations) |
| 166 | return variations |
| 167 | |
| 168 | |
| 169 | |