| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| import dask.config as dc |
| import dask.array as da |
| import torch |
| import numpy as np |
| import zarr |
| import pandas as pd |
| from datetime import datetime |
| import time |
| import os |
|
|
| from atmorep.datasets.normalizer import normalize |
| from atmorep.utils.utils import tokenize, get_weights |
|
|
| class MultifieldDataSampler( torch.utils.data.IterableDataset): |
| |
| |
| def __init__( self, file_path, fields, years, batch_size, pre_batch, n_size, |
| num_samples, with_shuffle = False, time_sampling = 1, with_source_idxs = False, compute_weights = False, |
| fields_targets = None, pre_batch_targets = None ) : |
| ''' |
| Data set for single dynamic field at an arbitrary number of vertical levels |
| |
| nsize : neighborhood in (tsteps, deg_lat, deg_lon) |
| ''' |
| super( MultifieldDataSampler).__init__() |
|
|
| self.fields = fields |
| self.batch_size = batch_size |
| self.n_size = n_size |
| self.num_samples = num_samples |
| self.with_source_idxs = with_source_idxs |
| self.compute_weights = compute_weights |
| self.with_shuffle = with_shuffle |
| self.pre_batch = pre_batch |
| |
| assert os.path.exists(file_path), f"File path {file_path} does not exist" |
| self.ds = zarr.open( file_path) |
| |
| self.dask_array_data = da.from_zarr(self.ds['data']) |
| self.dask_array_sfc = da.from_zarr(self.ds['data_sfc']) |
|
|
| self.ds_global = self.ds.attrs['is_global'] |
|
|
| self.lats = np.array( self.ds['lats']) |
| self.lons = np.array( self.ds['lons']) |
| |
| sh = self.ds['data'].shape |
| st = self.ds['time'].shape |
| self.ds_len = st[0] |
| print( f'self.ds[\'data\'] : {sh} :: {st}') |
| print( f'self.lats : {self.lats.shape}', flush=True) |
| print( f'self.lons : {self.lons.shape}', flush=True) |
| self.fields_idxs = [] |
|
|
| self.time_sampling = time_sampling |
| self.range_lat = np.array( self.lats[ [0,-1] ]) |
| self.range_lon = np.array( self.lons[ [0,-1] ]) |
| self.res = np.array(self.ds.attrs['res']) |
| self.year_base = self.ds['time'][0].astype(datetime).year |
|
|
| |
| self.range_lat += np.array([n_size[1] / 2., -n_size[1] / 2.]) |
| |
| if self.ds_global < 1.: |
| self.range_lon += np.array([n_size[2]/2., -n_size[2]/2.]) |
| |
| |
| self.normalizers = [] |
| for ifield, field_info in enumerate(fields) : |
| corr_type = 'global' if len(field_info) <= 6 else field_info[6] |
| nf_name = 'global_norm' if corr_type == 'global' else 'norm' |
| self.normalizers.append( [] ) |
| for vl in field_info[2]: |
| if vl == 0: |
| field_idx = self.ds.attrs['fields_sfc'].index( field_info[0]) |
| n_name = f'normalization/{nf_name}_sfc' |
| self.normalizers[ifield] += [self.ds[n_name].oindex[ :, :, field_idx]] |
| else: |
| vl_idx = self.ds.attrs['levels'].index(vl) |
| field_idx = self.ds.attrs['fields'].index( field_info[0]) |
| n_name = f'normalization/{nf_name}' |
| self.normalizers[ifield] += [self.ds[n_name].oindex[ :, :, field_idx, vl_idx]] |
| |
| |
| self.times = pd.DatetimeIndex( self.ds['time']) |
| idxs_years = self.times.year == years[0] |
| for year in years[1:] : |
| idxs_years = np.logical_or( idxs_years, self.times.year == year) |
| self.idxs_years = np.where( idxs_years)[0] |
|
|
| self.num_samples = min( self.num_samples, self.idxs_years.shape[0]) |
|
|
| |
| def shuffle( self) : |
|
|
| worker_info = torch.utils.data.get_worker_info() |
| rng_seed = None |
| if worker_info is not None : |
| rng_seed = int(time.time()) // (worker_info.id+1) + worker_info.id |
|
|
| rng = np.random.default_rng( rng_seed) |
| self.idxs_perm_t = rng.permutation( self.idxs_years)[ : self.num_samples // self.batch_size] |
| |
| lats = rng.random(self.num_samples) * (self.range_lat[1] - self.range_lat[0]) +self.range_lat[0] |
| lons = rng.random(self.num_samples) * (self.range_lon[1] - self.range_lon[0]) +self.range_lon[0] |
|
|
| |
| res_inv = 1.0 / self.res * 1.00001 |
| lats = self.res[0] * np.round( lats * res_inv[0]) |
| lons = self.res[1] * np.round( lons * res_inv[1]) |
|
|
| self.idxs_perm = np.stack( [lats, lons], axis=1) |
|
|
| |
| def __iter__(self): |
|
|
| if self.with_shuffle : |
| self.shuffle() |
|
|
| lats, lons = self.lats, self.lons |
| ts, n_size = self.time_sampling, self.n_size |
| ns_2 = np.array(self.n_size) / 2. |
| res = self.res |
|
|
| iter_start, iter_end = self.worker_workset() |
|
|
| for bidx in range( iter_start, iter_end) : |
|
|
| sources, token_infos = [[] for _ in self.fields], [[] for _ in self.fields] |
| sources_infos, source_idxs = [], [] |
| |
| i_bidx = self.idxs_perm_t[bidx] |
| idxs_t = list(np.arange( i_bidx - n_size[0]*ts, i_bidx, ts, dtype=np.int64)) |
| |
| |
| with dc.set(**{'array.slicing.split_large_chunks': True}): |
| data_tt_sfc = self.dask_array_sfc[idxs_t].compute() |
| data_tt = self.dask_array_data[idxs_t].compute() |
|
|
|
|
| for sidx in range(self.batch_size) : |
| |
| idx = self.idxs_perm[bidx*self.batch_size+sidx] |
| |
| lat_ran = np.where(np.logical_and(lats>idx[0]-ns_2[1]-res[0]/2.,lats<idx[0]+ns_2[1]))[0] |
| |
| assert not ((idx[1]-ns_2[2]) < 0. and (idx[1]+ns_2[2]) > 360.) |
| il, ir = (idx[1]-ns_2[2]-res[1]/2., idx[1]+ns_2[2]) |
| if il < 0. : |
| lon_ran = np.concatenate( [np.where( lons > il+360)[0], np.where(lons < ir)[0]], 0) |
| elif ir > 360. : |
| lon_ran = np.concatenate( [np.where( lons > il)[0], np.where(lons < ir-360)[0]], 0) |
| else : |
| lon_ran = np.where(np.logical_and( lons > il, lons < ir))[0] |
| |
| sources_infos += [ [ self.ds['time'][ idxs_t ].astype(datetime), |
| self.lats[lat_ran], self.lons[lon_ran], self.res ] ] |
|
|
| if self.with_source_idxs : |
| source_idxs += [ (idxs_t, lat_ran, lon_ran) ] |
|
|
| |
| for ifield, field_info in enumerate(self.fields): |
| source_lvl, tok_info_lvl = [], [] |
| tok_size = field_info[4] |
| num_tokens = field_info[3] |
| corr_type = 'global' if len(field_info) <= 6 else field_info[6] |
| |
| for ilevel, vl in enumerate(field_info[2]): |
| if vl == 0 : |
| field_idx = self.ds.attrs['fields_sfc'].index( field_info[0]) |
| data_t = data_tt_sfc[ :, field_idx ] |
| else : |
| field_idx = self.ds.attrs['fields'].index( field_info[0]) |
| vl_idx = self.ds.attrs['levels'].index(vl) |
| data_t = data_tt[ :, field_idx, vl_idx ] |
| |
| source_data, tok_info = [], [] |
| |
| cdata = data_t[ ... , lat_ran[:,np.newaxis], lon_ran[np.newaxis,:]] |
| |
| normalizer = self.normalizers[ifield][ilevel] |
|
|
| if corr_type != 'global': |
| |
| if lat_ran[0] < lat_ran[-1] and lon_ran[0] < lon_ran[-1]: |
| lat_max, lat_min = max(lat_ran), min(lat_ran) |
| lon_max, lon_min = max(lon_ran), min(lon_ran) |
| normalizer = normalizer[:,:,lat_min:lat_max+1,lon_min:lon_max+1] |
| |
| |
| else: |
| normalizer = normalizer[ ... , lat_ran[:,np.newaxis], lon_ran[np.newaxis,:]] |
| |
| |
| cdata = normalize(cdata, normalizer, sources_infos[-1][0], year_base = self.year_base) |
| |
| source_data = tokenize( torch.from_numpy( cdata), tok_size ) |
| |
| dates = self.ds['time'][ idxs_t ].astype(datetime) |
| cdates = dates[tok_size[0]-1::tok_size[0]] |
| |
| dates = [(d.year, d.timetuple().tm_yday-1, d.hour) for d in cdates] |
| lats_sidx = self.lats[lat_ran][ tok_size[1]//2 :: tok_size[1] ] |
| lons_sidx = self.lons[lon_ran][ tok_size[2]//2 :: tok_size[2] ] |
| |
| tok_info += [[[[[ year, day, hour, vl, lat, lon, vl, self.res[0]] for lon in lons_sidx] |
| for lat in lats_sidx] |
| for (year, day, hour) in dates]] |
|
|
| source_lvl += [ source_data ] |
| tok_info_lvl += [ torch.tensor(tok_info, dtype=torch.float32).flatten( 1, -2)] |
| sources[ifield] += [ torch.stack(source_lvl, 0) ] |
| token_infos[ifield] += [ torch.stack(tok_info_lvl, 0) ] |
| |
| |
| sources = [torch.stack(sources_field).transpose(1,0) for sources_field in sources] |
| token_infos = [torch.stack(tis_field).transpose(1,0) for tis_field in token_infos] |
| sources = self.pre_batch( sources, token_infos ) |
|
|
| tmidx_list = sources[-1] |
| weights_idx_list = [] |
| if self.compute_weights: |
| for ifield, field_info in enumerate(self.fields): |
| weights = [] |
| for ilevel, vl in enumerate(field_info[2]): |
| for ibatch in range(self.batch_size): |
| |
| lats_idx = source_idxs[ibatch][1] |
| lons_idx = source_idxs[ibatch][2] |
|
|
| idx_base = tmidx_list[ifield][ilevel][ibatch] |
| idx_loc = idx_base - np.prod(num_tokens) * ibatch |
| |
| grid = np.flip(np.array( np.meshgrid( lons_idx, lats_idx)), axis = 0) |
| grid = torch.from_numpy( np.array( np.broadcast_to( grid, |
| shape = [tok_size[0]*num_tokens[0], *grid.shape])).swapaxes(0,1)) |
|
|
| grid_lats_toked = tokenize( grid[0], tok_size).flatten( 0, 2) |
|
|
| lats_mskd_b = np.array([np.unique(t) for t in grid_lats_toked[ idx_loc ].numpy()]) |
|
|
| weights.append([get_weights(la) for la in lats_mskd_b]) |
|
|
| weights_idx_list.append(weights) |
| sources = (*sources, weights_idx_list) |
|
|
| |
| targets, target_info = None, None |
| target_idxs = None |
| |
| yield ( sources, targets, (source_idxs, sources_infos), (target_idxs, target_info)) |
|
|
| |
| def set_data( self, times_pos, batch_size = None) : |
| ''' |
| times_pos = np.array( [ [year, month, day, hour, lat, lon], ...] ) |
| - lat \in [90,-90] = [90N, 90S] |
| - lon \in [0,360] |
| - (year,month) pairs should be a limited number since all data for these is loaded |
| ''' |
| |
| self.idxs_perm = np.zeros( (len(times_pos), 2)) |
| self.idxs_perm_t = [] |
| self.num_samples = len(times_pos) |
| for idx, item in enumerate( times_pos) : |
|
|
| assert item[2] >= 1 and item[2] <= 31 |
| assert item[3] >= 0 and item[3] < int(24 / self.time_sampling) |
| assert item[4] >= -90. and item[4] <= 90. |
|
|
| tstamp = pd.to_datetime( f'{item[0]}-{item[1]}-{item[2]}-{item[3]}', format='%Y-%m-%d-%H') |
| |
| self.idxs_perm_t += [ np.where( self.times == tstamp)[0]+1 ] |
|
|
| |
| self.idxs_perm[idx] = np.array( [90. - item[4], item[5]]) |
| |
| self.idxs_perm_t = np.array(self.idxs_perm_t).squeeze() |
|
|
| |
| def set_global( self, times, batch_size = None, token_overlap = [0, 0]) : |
| ''' generate patch/token positions for global grid ''' |
| token_overlap = np.array( token_overlap).astype(np.int64) |
|
|
| |
| ifield = 0 |
| field = self.fields[ifield] |
|
|
| res = self.res |
| side_len = np.array( [field[3][1] * field[4][1]*res[0], field[3][2] * field[4][2]*res[1]] ) |
| overlap = np.array([token_overlap[0]*field[4][1]*res[0],token_overlap[1]*field[4][2]*res[1]]) |
| side_len_2 = side_len / 2. |
| assert all( overlap <= side_len_2), 'token_overlap too large for #tokens, reduce if possible' |
|
|
| |
| times_pos = [] |
| for ctime in times : |
|
|
| lat = side_len_2[0].item() |
| num_tiles_lat = 0 |
| while (lat + side_len_2[0].item()) < 180. : |
| num_tiles_lat += 1 |
| lon = side_len_2[1].item() - overlap[1].item()/2. |
| num_tiles_lon = 0 |
| while (lon - side_len_2[1]) < 360. : |
| times_pos += [[*ctime, -lat + 90., np.mod(lon,360.) ]] |
| lon += side_len[1].item() - overlap[1].item() |
| num_tiles_lon += 1 |
| lat += side_len[0].item() - overlap[0].item() |
|
|
| |
| |
| |
| |
| lat -= side_len[0] - overlap[0] |
| if lat - side_len_2[0] < 180. : |
| num_tiles_lat += 1 |
| lat = 180. - side_len_2[0].item() + res[0] |
| lon = side_len_2[1].item() - overlap[1].item()/2. |
| while (lon - side_len_2[1]) < 360. : |
| times_pos += [[*ctime, -lat + 90., np.mod(lon,360.) ]] |
| lon += side_len[1].item() - overlap[1].item() |
|
|
| |
| batch_size = len(times_pos) |
| |
| print( 'Number of batches per global forecast: {}'.format( num_tiles_lat) ) |
|
|
| self.set_data( times_pos, batch_size) |
|
|
| |
| def __len__(self): |
| return self.num_samples // self.batch_size |
|
|
| |
| def worker_workset( self) : |
|
|
| worker_info = torch.utils.data.get_worker_info() |
|
|
| if worker_info is None: |
| iter_start = 0 |
| iter_end = self.num_samples |
| |
| else: |
| |
| per_worker = len(self) // worker_info.num_workers |
| worker_id = worker_info.id |
| iter_start = int(worker_id * per_worker) |
| iter_end = int(iter_start + per_worker) |
| if worker_info.id+1 == worker_info.num_workers : |
| iter_end = len(self) |
|
|
| return iter_start, iter_end |
|
|
|
|