(args)
| 109 | return args |
| 110 | |
| 111 | def fit(args): |
| 112 | positional_encoding_name = args.pe |
| 113 | neural_network_name = args.nn |
| 114 | dataset = args.dataset |
| 115 | |
| 116 | torch.manual_seed(args.seed) |
| 117 | np.random.seed(args.seed) |
| 118 | random.seed(args.seed) |
| 119 | |
| 120 | with open(args.hparams) as f: |
| 121 | hparams = yaml.safe_load(f) |
| 122 | |
| 123 | dataset_hparams = hparams[dataset]["dataset"] |
| 124 | |
| 125 | hparams = hparams[dataset] |
| 126 | print(args) |
| 127 | if args.use_expnamehps: |
| 128 | if 'seed' in args.expname: |
| 129 | appender_in_yaml = args.expname.split('_seed')[0] |
| 130 | else: |
| 131 | appender_in_yaml = args.expname |
| 132 | hparams = hparams[f"{positional_encoding_name}-{neural_network_name}-{appender_in_yaml}"] |
| 133 | else: |
| 134 | hparams = hparams[f"{positional_encoding_name}-{neural_network_name}"] |
| 135 | hparams.update(dataset_hparams) |
| 136 | |
| 137 | hparams = overwrite_hparams_with_args(hparams, args) |
| 138 | hparams = set_default_if_unset(hparams, "max_radius", 360) |
| 139 | |
| 140 | if args.dataset == "landoceandataset": |
| 141 | datamodule = LandOceanDataModule(batch_size=hparams["batch_size"]) |
| 142 | elif args.dataset == "inat2018": |
| 143 | datamodule = Inat2018DataModule(hparams["inat_directory"], batch_size=hparams["batch_size"], mode="location") |
| 144 | elif args.dataset == "checkerboard": |
| 145 | datamodule = CheckerboardDataModule(num_samples=hparams["num_samples"], |
| 146 | num_classes=hparams["num_classes"], |
| 147 | num_support=int(hparams["num_support"] * args.checkerboard_scale), |
| 148 | batch_size=hparams["batch_size"]) |
| 149 | elif args.dataset == 'era5dataset': |
| 150 | datamodule = ERA5DataModule(batch_size=hparams["batch_size"], |
| 151 | data_root=hparams["era5_directory"], |
| 152 | label_key="t2m") |
| 153 | elif args.dataset == 'era5dataset_multi': |
| 154 | datamodule = ERA5DataModule(batch_size=hparams["batch_size"], |
| 155 | data_root=hparams["era5_directory"], |
| 156 | label_key=['u10', 'v10', 't2m', 'sp', 'd2m', 'ssr', 'str', 'tp']) |
| 157 | |
| 158 | if args.resume_ckpt_from_results_dir: |
| 159 | resume_checkpoint = find_best_checkpoint(parse_resultsdir(args), |
| 160 | f"{positional_encoding_name}-{neural_network_name}", |
| 161 | verbose=True) |
| 162 | else: |
| 163 | resume_checkpoint = None |
| 164 | |
| 165 | locationencoder = LocationEncoder( |
| 166 | positional_encoding_name, |
| 167 | neural_network_name, |
| 168 | hparams=hparams |
no test coverage detected