MCPcopy Create free account
hub / github.com/MarcCoru/locationencoder / fit

Function fit

train.py:111–358  ·  view source on GitHub ↗
(args)

Source from the content-addressed store, hash-verified

109 return args
110
111def 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

Callers 6

fit_comparisonFunction · 0.90
fit_shFunction · 0.90
mainFunction · 0.90
fit_modelsFunction · 0.90
fit_modelsFunction · 0.90
train.pyFile · 0.70

Calls 15

test_dataloaderMethod · 0.95
train_dataloaderMethod · 0.95
get_test_locsMethod · 0.95
set_default_if_unsetFunction · 0.90
LandOceanDataModuleClass · 0.90
Inat2018DataModuleClass · 0.90
ERA5DataModuleClass · 0.90
find_best_checkpointFunction · 0.90
parse_resultsdirFunction · 0.90
LocationEncoderClass · 0.90
count_parametersFunction · 0.90

Tested by

no test coverage detected