import os import torch import numpy as np import xarray as xr from torch.utils.data import Dataset import glob class NDVITimeSeriesDataset(Dataset): def __init__(self, sequence_length=3, spatial=False): """ Đọc dữ liệu S2 từ cache, tính NDVI và tạo Time-Series. spatial=False -> Output 1D cho LSTM/ARIMA spatial=True -> Output 2D cho ConvLSTM """ self.sequence_length = sequence_length self.spatial = spatial self.data_seqs = [] self.targets = [] # Load from cache cache_files = glob.glob("dataset_cache/*.nc") s2_files = [f for f in cache_files if len(os.path.basename(f)) == 35] # S2 cache filenames usually have length 32 + 3 (.nc) if not s2_files: print("[WARNING] Không tìm thấy dữ liệu S2 trong cache! Dùng dummy data.") self._create_dummy() return try: print(f"[DATA] Loading real data from {s2_files[0]}") ds = xr.open_dataset(s2_files[0], engine='netcdf4') if 'time' not in ds.dims or len(ds.time) < sequence_length + 1: self._create_dummy() return # Tính NDVI: (B08 - B04) / (B08 + B04) b8 = ds['B08'].astype(np.float32) b4 = ds['B04'].astype(np.float32) ndvi = (b8 - b4) / (b8 + b4 + 1e-8) ndvi = ndvi.fillna(0).values # shape: (time, y, x) # Lấy 1 pixel trung tâm hoặc toàn bộ ảnh if not self.spatial: # Average pooling over space for 1D time series ndvi = ndvi.mean(axis=(1, 2)) # shape: (time,) for i in range(len(ndvi) - sequence_length): self.data_seqs.append(ndvi[i:i+sequence_length]) self.targets.append(ndvi[i+sequence_length]) else: # Spatial data for ConvLSTM # Downsample to 64x64 to avoid OOM from skimage.transform import resize T = len(ndvi) ndvi_resized = np.zeros((T, 64, 64)) for t in range(T): ndvi_resized[t] = resize(ndvi[t], (64, 64)) for i in range(T - sequence_length): self.data_seqs.append(ndvi_resized[i:i+sequence_length]) # (seq, 64, 64) self.targets.append(ndvi_resized[i+sequence_length]) # (64, 64) except Exception as e: print(f"[ERROR] {e}. Dùng dummy data.") self._create_dummy() def _create_dummy(self): T = 20 if not self.spatial: ndvi = np.random.rand(T).astype(np.float32) for i in range(T - self.sequence_length): self.data_seqs.append(ndvi[i:i+self.sequence_length]) self.targets.append(ndvi[i+self.sequence_length]) else: ndvi = np.random.rand(T, 64, 64).astype(np.float32) for i in range(T - self.sequence_length): self.data_seqs.append(ndvi[i:i+self.sequence_length]) self.targets.append(ndvi[i+self.sequence_length]) def __len__(self): return len(self.data_seqs) def __getitem__(self, idx): x = torch.tensor(self.data_seqs[idx], dtype=torch.float32) y = torch.tensor(self.targets[idx], dtype=torch.float32) if not self.spatial: x = x.unsqueeze(1) # (seq_len, features=1) y = y.unsqueeze(0) # (1,) else: x = x.unsqueeze(1) # (seq_len, channels=1, H, W) y = y.unsqueeze(0) # (1, H, W) return x, y