Setup linear classifiers with different configurations.
(
sample_output,
n_last_blocks_list: Tuple[int, ...],
learning_rates: Tuple[float, ...],
batch_size: int,
num_classes: int = 1000,
device: torch.device = None,
)
| 219 | |
| 220 | |
| 221 | def setup_linear_classifiers( |
| 222 | sample_output, |
| 223 | n_last_blocks_list: Tuple[int, ...], |
| 224 | learning_rates: Tuple[float, ...], |
| 225 | batch_size: int, |
| 226 | num_classes: int = 1000, |
| 227 | device: torch.device = None, |
| 228 | ) -> Tuple[AllClassifiers, List[Dict]]: |
| 229 | """Setup linear classifiers with different configurations.""" |
| 230 | linear_classifiers_dict = nn.ModuleDict() |
| 231 | optim_param_groups = [] |
| 232 | |
| 233 | for n in n_last_blocks_list: |
| 234 | for avgpool in [True]: |
| 235 | for _lr in learning_rates: |
| 236 | lr = scale_lr(_lr, batch_size) |
| 237 | out_dim = create_linear_input(sample_output, use_n_blocks=n, use_avgpool=avgpool).shape[1] |
| 238 | linear_classifier = LinearClassifier( |
| 239 | out_dim, use_n_blocks=n, use_avgpool=avgpool, num_classes=num_classes |
| 240 | ) |
| 241 | linear_classifier = linear_classifier.to(device) |
| 242 | classifier_key = f"classifier_{n}_blocks_avgpool_{avgpool}_lr_{lr:.5f}".replace(".", "_") |
| 243 | if is_main_process(): |
| 244 | logger.info(f"Create linear classifier {classifier_key} with input_dim={out_dim}") |
| 245 | linear_classifiers_dict[classifier_key] = linear_classifier |
| 246 | optim_param_groups.append({"params": linear_classifier.parameters(), "lr": lr}) |
| 247 | |
| 248 | linear_classifiers = AllClassifiers(linear_classifiers_dict) |
| 249 | if dist.is_initialized(): |
| 250 | linear_classifiers = nn.parallel.DistributedDataParallel( |
| 251 | linear_classifiers, device_ids=[get_rank() % torch.cuda.device_count()] |
| 252 | ) |
| 253 | |
| 254 | return linear_classifiers, optim_param_groups |
| 255 | |
| 256 | |
| 257 | def train_one_epoch( |
no test coverage detected