11from __future__ import annotations
22
3+ import importlib
4+ import sys
35from pathlib import Path
46from typing import Any
57
68import numpy as np
7- import rasterio as rio
89import torch
910
1011from opensr_srgan .data .utils .normalizer import Normalizer
@@ -27,9 +28,7 @@ def __init__(self, config: Any, phase=None, taco_file=None):
2728 raise ValueError ("SEN2NAIP requires config.Data." )
2829
2930 if taco_file is None :
30- taco_file = getattr (
31- data_cfg , "sen2naip_taco_file" , DEFAULT_SEN2NAIP_TACO_FILE
32- )
31+ taco_file = getattr (data_cfg , "sen2naip_taco_file" , None )
3332 if phase is None :
3433 phase = getattr (data_cfg , "sen2naip_phase" , "train" )
3534 val_fraction = getattr (data_cfg , "sen2naip_val_fraction" , 0.1 )
@@ -78,6 +77,10 @@ def _to_tensor(data: np.ndarray) -> torch.Tensor:
7877 return tensor
7978
8079 def __getitem__ (self , idx ):
80+ rio = sys .modules .get ("rasterio" )
81+ if rio is None :
82+ rio = importlib .import_module ("rasterio" )
83+
8184 sample = self .dataset .read (self .indices [idx ])
8285 lr_path = sample .read (0 )
8386 hr_path = sample .read (1 )
@@ -96,9 +99,8 @@ def __getitem__(self, idx):
9699
97100
98101if __name__ == "__main__" :
99- DEFAULT_SEN2NAIP_TACO_FILE = "/data1/datasets/SEN2NAIP/sen2naipv2-crosssensor.taco"
100102 ds = SEN2NAIP (
101103 config = "opensr_srgan/configs/config_10m.yaml" ,
102104 phase = "train" ,
103- taco_file = DEFAULT_SEN2NAIP_TACO_FILE ,
105+ taco_file = "/data1/datasets/SEN2NAIP/sen2naipv2-crosssensor.taco" ,
104106 )
0 commit comments