94 lines
3.7 KiB
Python
94 lines
3.7 KiB
Python
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
|