Skip to content

Commit f478106

Browse files
committed
tests for PL
1 parent 43aeb26 commit f478106

1 file changed

Lines changed: 8 additions & 6 deletions

File tree

opensr_srgan/data/sen2naip/sen2naip_dataset.py

Lines changed: 8 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -1,10 +1,11 @@
11
from __future__ import annotations
22

3+
import importlib
4+
import sys
35
from pathlib import Path
46
from typing import Any
57

68
import numpy as np
7-
import rasterio as rio
89
import torch
910

1011
from 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

98101
if __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

Comments
 (0)