#%%单个模式预估计算
import xarray as xr
import numpy as np
import os
import gc
from tqdm import tqdm
from glob import glob
import pandas as pd
import warnings
warnings.filterwarnings("ignore", category=xr.SerializationWarning)
basepath = 'E://'
os.chdir(basepath)
# 定义阈值
snow_threshold = 2.0  # mm/day
rain_threshold = 2.0  # mm/day
snc_threshold = 50.0  # %
       
model_list = ['CanESM5','CESM2','CMCC-CM2-SR5','CMCC-ESM2','CNRM-CM6-1',
              'CNRM-ESM2-1','EC-Earth3','EC-Earth3-CC','EC-Earth3-Veg','EC-Earth3-Veg-LR',
              'INM-CM4-8','INM-CM5-0','MIROC6','MIROC-ES2L','MIROC-ES2H',
              'MPI-ESM1-2-HR','MPI-ESM1-2-LR','MRI-ESM2-0','UKESM1-0-LL','GFDL-CM4',
              'AWI-ESM-1-REcoM','HadGEM3-GC31-LL']  # 替换为实际模式名称
# 定义时间段
time_periods = {
    'hist': slice('2005', '2014'),
    'ssp245': slice('2091', '2100'),
    'ssp585': slice('2091', '2100')
}
# 处理每个模式
for model in tqdm(model_list):
    print(f"Processing model: {model}")
    
    # 存储每个实验的结果
    results = {}
    
    # 处理每个实验
    for experiment in ['hist', 'ssp245', 'ssp585']:
        print(f"  Processing experiment: {experiment}")
        
        def get_files(base, model, experiment, var):
            return sorted(
                glob(f"{base}/{model}/{experiment}/ensemble1/{var}_*.nc") +
                glob(f"{base}/{model}/{experiment}/ensemble1/{var}_*.hdf")
            )

        pr_files = get_files("precipitation_flux", model, experiment, "pr")
        prsn_files = get_files("snowfall_flux", model, experiment, "prsn")
        snc_files = get_files("snow_area_fraction", model, experiment, "snc")
        # 查找文件
        #pr_files = glob(pr_path)
       # prsn_files = glob(prsn_path)
        #snc_files = glob(snc_path)
        
        if not pr_files:
            print("    No pr files found at pr")
            continue
        if not prsn_files:
            print("    No prsn files found at prsn")
            continue
        if not snc_files:
            print("    No snc files found at snc")
            continue
        
        print(f"    Found {len(pr_files)} pr files, {len(prsn_files)} prsn files, {len(snc_files)} snc files")
        
        try:
            # 使用xarray打开数据集
            print("    Opening datasets...")
            pr_data = xr.open_mfdataset(pr_files)
            prsn_data = xr.open_mfdataset(prsn_files)
            snc_data = xr.open_mfdataset(snc_files)
            
            # 选择时间段
            print(f"    Selecting time period {time_periods[experiment]}...")
            pr_data = pr_data.sel(time=time_periods[experiment])
            prsn_data = prsn_data.sel(time=time_periods[experiment])
            snc_data = snc_data.sel(time=time_periods[experiment])
            
            # 确保所有数据集具有相同的时间维度
            print("    Aligning time dimensions...")
            common_times = np.intersect1d(pr_data.time.values, 
                                        np.intersect1d(prsn_data.time.values, snc_data.time.values))
            
            if len(common_times) == 0:
                print(f"    No common time steps found, skipping {model}/{experiment}")
                # 清理已打开的数据集
                pr_data.close()
                prsn_data.close()
                snc_data.close()
                continue
            pr_data = pr_data.sel(time=common_times)
            prsn_data = prsn_data.sel(time=common_times)
            snc_data = snc_data.sel(time=common_times)
            
            # 单位转换：从 kg/m²/s 到 mm/day，乘以86400(秒/天)
            print("    Converting units from kg/m²/s to mm/day...")
            pr_mm = pr_data['pr'] * 86400
            prsn_mm = prsn_data['prsn'] * 86400
            
            # 计算降雨量 (pr - prsn) 单位已转换为mm/day
            print("    Calculating rain from pr and prsn...")
            rain_mm = pr_mm - prsn_mm
            
            # 根据阈值计算每日的snow_days, rain_days, ros_days (布尔型标记)
            print("    Identifying days according to thresholds...")
            snow_days_binary = (prsn_mm >= snow_threshold).astype(int)
            rain_days_binary = (rain_mm >= rain_threshold).astype(int)
            ros_days_binary = rain_days_binary * (snc_data['snc'] >= snc_threshold).astype(int)
            
            # 计算每年的日数 - 使用时间的年份而不是添加新坐标
            print("    Calculating yearly day counts...")
            # 创建年份数组用于分组
            snow_days_binary = snow_days_binary.assign_coords(year=snow_days_binary['time'].dt.year)
            rain_days_binary = rain_days_binary.assign_coords(year=rain_days_binary['time'].dt.year)
            ros_days_binary = ros_days_binary.assign_coords(year=ros_days_binary['time'].dt.year)

            
            # 计算每年的总天数
            yearly_snow_days = snow_days_binary.groupby('year').sum(dim='time')
            yearly_rain_days = rain_days_binary.groupby('year').sum(dim='time')
            yearly_ros_days = ros_days_binary.groupby('year').sum(dim='time')
            
            # 计算10年平均值
            print("    Calculating 10-year averages of yearly counts...")
            snow_days_mean = yearly_snow_days.mean(dim='year')
            rain_days_mean = yearly_rain_days.mean(dim='year')
            ros_days_mean = yearly_ros_days.mean(dim='year')
            
            # 计算10年平均的雪覆盖
            snc_mean = snc_data['snc'].mean(dim='time')
            # 雪覆盖边界 (snc ≥ 50%)
            snc_border = (snc_mean >= snc_threshold).astype(int)
            
            # 存储每日数据 - 确保计算正确并转换为实际数据，避免nan
            # 这里我们复制一下数据，确保它是独立的且已经计算
            print("    Storing daily data...")
            snow_days_daily = snow_days_binary.compute().copy(deep=True)
            rain_days_daily = rain_days_binary.compute().copy(deep=True)
            ros_days_daily = ros_days_binary.compute().copy(deep=True)
            # 为每个实验设置正确的时间轴
            snow_days_daily = snow_days_daily.assign_coords(time=pr_data['time'].sel(time=common_times))
            rain_days_daily = rain_days_daily.assign_coords(time=pr_data['time'].sel(time=common_times))
            ros_days_daily = ros_days_daily.assign_coords(time=pr_data['time'].sel(time=common_times))

            # 存储结果 - 每天的情况和10年平均值
            print("    Storing results for {experiment}...")
            results[experiment] = {
                'snow_days': snow_days_mean,  # 10年平均的年雪日数
                'rain_days': rain_days_mean,  # 10年平均的年雨日数
                'ros_days': ros_days_mean,    # 10年平均的年雨雪并存日数
                'snc_border': snc_border,     # 10年平均的雪覆盖边界
                'snow_days_daily': snow_days_daily,  # 每天的雪日情况
                'rain_days_daily': rain_days_daily,  # 每天的雨日情况
                'ros_days_daily': ros_days_daily     # 每天的雨雪并存日情况
            }
            
            # 关闭数据集，释放内存
            print("    Closing datasets...")
            pr_data.close()
            prsn_data.close()
            snc_data.close()
            
            # 强制垃圾回收
            gc.collect()
            
        except Exception as e:
            print(f"    Error processing {model}/{experiment}: {str(e)}")
            # 尝试关闭可能已打开的数据集
            try:
                if 'pr_data' in locals():
                    pr_data.close()
                if 'prsn_data' in locals():
                    prsn_data.close()
                if 'snc_data' in locals():
                    snc_data.close()
            except:
                pass
            gc.collect()
            continue
    
    # 确保历史数据存在
    if 'hist' not in results:
        print(f"No historical data for {model}, skipping change calculations.")
        continue
    
    # 计算相对历史差值
    print("  Calculating changes relative to historical period...")
    for experiment in ['ssp245', 'ssp585']:
        if experiment not in results:
            print(f"  No {experiment} data for {model}, skipping this scenario.")
            continue
            
        # 计算变化量 - 年度日数的10年平均值之差
        print(f"    Calculating changes for {experiment}...")
        results[f'{experiment}_change'] = {
            'snow_days_change': results[experiment]['snow_days'] - results['hist']['snow_days'],
            'rain_days_change': results[experiment]['rain_days'] - results['hist']['rain_days'],
            'ros_days_change': results[experiment]['ros_days'] - results['hist']['ros_days']
        }
    
    # 保存 hist 数据
    if 'hist' in results:
        try:
            hist_ds = xr.Dataset()
            for var_name, var_data in results['hist'].items():
                hist_ds[f"hist_{var_name}"] = var_data.compute() if hasattr(var_data, 'compute') else var_data
            hist_file = f"/single_models//{model}_hist.nc"
            print(f"  Saving hist results to {hist_file}...")
            hist_ds.to_netcdf(hist_file)
            hist_ds.close()
            del hist_ds
        except Exception as e:
            print(f"  Error saving hist results for {model}: {str(e)}")

    # 保存 ssp245、ssp585 和 change 数据（合并为一个 future 文件）
    try:
        future_ds = xr.Dataset()
        for exp in ['ssp245', 'ssp585']:
            if exp in results:
                for var_name, var_data in results[exp].items():
                    future_ds[f"{exp}_{var_name}"] = var_data.compute() if hasattr(var_data, 'compute') else var_data

            change_key = f"{exp}_change"
            if change_key in results:
                for var_name, var_data in results[change_key].items():
                    future_ds[f"{change_key}_{var_name}"] = var_data.compute() if hasattr(var_data, 'compute') else var_data

        future_file = f"/single_models//{model}_projection.nc"
        print(f"  Saving future results to {future_file}...")
        future_ds.to_netcdf(future_file)
        future_ds.close()
        del future_ds
    except Exception as e:
        print(f"  Error saving future results for {model}: {str(e)}")
        
    print("="*50 + f"{model}")
    
print("All models processed successfully.")
#%%创建洪涝模式集资料
import xarray as xr
import numpy as np
import os
import gc
from tqdm import tqdm
from glob import glob
import pandas as pd
import warnings
from functools import reduce  # 添加reduce函数导入
warnings.filterwarnings("ignore", category=xr.SerializationWarning)
warnings.filterwarnings("ignore", category=UserWarning)
basepath = 'E://'
os.chdir(basepath)
# 在代码开头添加
import dask

# 设置默认分块大小
dask.config.set({"array.chunk-size": "256MiB"})
#'CanESM5',
model_list = ['CESM2','CMCC-CM2-SR5','CMCC-ESM2','CNRM-CM6-1','CNRM-ESM2-1','EC-Earth3','EC-Earth3-CC','EC-Earth3-Veg','EC-Earth3-Veg-LR',       
              'INM-CM4-8','INM-CM5-0','MIROC6','MIROC-ES2L','MIROC-ES2H','MPI-ESM1-2-HR','MPI-ESM1-2-LR','MRI-ESM2-0','GFDL-CM4',
              'AWI-ESM-1-REcoM']  # 替换为实际模式名称

leap_year_models = [
    'CNRM-CM6-1', 'CNRM-ESM2-1', 'EC-Earth3-Veg-LR', 'EC-Earth3-CC',
    'EC-Earth3-Veg', 'MIROC6', 'EC-Earth3', 'MIROC-ES2L', 
    'MIROC-ES2H', 'MPI-ESM1-2-LR', 'MRI-ESM2-0', 'MPI-ESM1-2-HR',
    'AWI-ESM-1-REcoM'
]  # 需要调整闰年

# 定义时间段
time_periods = {
    'hist': slice('2005', '2014'),
    'ssp245': slice('2091', '2100'),
    'ssp585': slice('2091', '2100')
}

# 定义函数用于获取文件
def get_files(base, model, experiment, var):
    return sorted(
        glob(f"{base}/{model}/{experiment}/ensemble1/{var}_*.nc") +
        glob(f"{base}/{model}/{experiment}/ensemble1/{var}_*.hdf")
    )

# 定义函数用于处理单个实验
def process_experiment(model, experiment):
    print(f"  Processing experiment: {experiment}")
    
    pr_files = get_files("precipitation_flux", model, experiment, "pr")
    prsn_files = get_files("snowfall_flux", model, experiment, "prsn")
    snw_files = get_files("surface_snow_amount", model, experiment, "snw")
    tas_files = get_files("air_temperature", model, experiment, "tas")
    
    # 检查文件是否存在
    if not pr_files:
        print("    No pr files found")
        return None
    if not prsn_files:
        print("    No prsn files found")
        return None
    if not snw_files:
        print("    No snw files found")
        return None
    if not tas_files:
        print("    No tas files found")
        return None
    
    print(f"    Found {len(pr_files)} pr files, {len(prsn_files)} prsn files, {len(snw_files)} snw files, {len(tas_files)} tas files")
    
    try:
        # 使用xarray打开数据集
        print("    Opening datasets...")
        pr_data = xr.open_mfdataset(pr_files, chunks={'time': 50, 'lat': 64, 'lon': 64})
        prsn_data = xr.open_mfdataset(prsn_files, chunks={'time': 50, 'lat': 64, 'lon': 64})
        snw_data = xr.open_mfdataset(snw_files, chunks={'time': 50, 'lat': 64, 'lon': 64})
        tas_data = xr.open_mfdataset(tas_files, chunks={'time': 50, 'lat': 64, 'lon': 64})

        # 在插值前减少精度
        pr_data = pr_data.astype('float32')
        prsn_data = prsn_data.astype('float32')
        snw_data = snw_data.astype('float32')
        tas_data = tas_data.astype('float32')
        
        # 选择时间段
        print(f"    Selecting time period {time_periods[experiment]}...")
        pr_data = pr_data.sel(time=time_periods[experiment])
        prsn_data = prsn_data.sel(time=time_periods[experiment])
        snw_data = snw_data.sel(time=time_periods[experiment])
        tas_data = tas_data.sel(time=time_periods[experiment])
        
        # 确保所有数据集具有相同的时间维度
        print("    Aligning time dimensions...")
        # 分步骤计算交集，避免嵌套调用
        times_intersection1 = np.intersect1d(pr_data.time.values, prsn_data.time.values)
        times_intersection2 = np.intersect1d(times_intersection1, snw_data.time.values)
        common_times = np.intersect1d(times_intersection2, tas_data.time.values)
        
        if len(common_times) == 0:
            print(f"    No common time steps found, skipping {model}/{experiment}")
            # 清理已打开的数据集
            pr_data.close()
            prsn_data.close()
            snw_data.close()
            tas_data.close()
            return None
        
        pr_data = pr_data.sel(time=common_times)
        prsn_data = prsn_data.sel(time=common_times)
        snw_data = snw_data.sel(time=common_times)
        tas_data = tas_data.sel(time=common_times)
        
        # 执行时间插值或调整（如果需要）
        start_year = int(time_periods[experiment].start)
        end_year = int(time_periods[experiment].stop)
        
        # 对闰年模型移除2月29日
        if model in leap_year_models:
            print(f"    Adjusting leap years for {model}")
            pr_data = pr_data.sel(time=~((pr_data.time.dt.month == 2) & (pr_data.time.dt.day == 29)))
            prsn_data = prsn_data.sel(time=~((prsn_data.time.dt.month == 2) & (prsn_data.time.dt.day == 29)))
            snw_data = snw_data.sel(time=~((snw_data.time.dt.month == 2) & (snw_data.time.dt.day == 29)))
            tas_data = tas_data.sel(time=~((tas_data.time.dt.month == 2) & (tas_data.time.dt.day == 29)))
            
        # 空间插值到256x512网格
        print("    Performing spatial interpolation to 256x512 grid...")
        # 定义目标网格
        target_lat = np.linspace(-90, 90, 256)
        target_lon = np.linspace(0, 360, 512, endpoint=False)
        
        # 对每个数据集执行插值
        pr_data = pr_data.interp(lat=target_lat, lon=target_lon, method='linear')
        prsn_data = prsn_data.interp(lat=target_lat, lon=target_lon, method='linear')
        snw_data = snw_data.interp(lat=target_lat, lon=target_lon, method='linear')
        tas_data = tas_data.interp(lat=target_lat, lon=target_lon, method='linear')
        
        # 单位转换：从 kg/m²/s 到 mm/day，乘以86400(秒/天)
        print("    Converting units from kg/m²/s to mm/day...")
        pr_mm = pr_data['pr'] * 86400
        prsn_mm = prsn_data['prsn'] * 86400
        tas = tas_data['tas']
        snw = snw_data['snw']
        
        # 检查并处理降雪量大于总降水量的情况
        invalid_mask = prsn_mm > pr_mm
        if invalid_mask.any():
            print(f"    Found {invalid_mask.sum().values} points where prsn > pr, setting both to NaN")
            pr_mm = pr_mm.where(~invalid_mask)
            prsn_mm = prsn_mm.where(~invalid_mask)
        
        # 计算降雨量 (pr - prsn) 单位已转换为mm/day
        print("    Calculating rain from pr and prsn...")
        rain_mm = pr_mm - prsn_mm
        
        
        
        # 创建结果数据集
        result_ds = xr.Dataset({
            f"{experiment}_rn": rain_mm,
            f"{experiment}_tas": tas,
            f"{experiment}_snw": snw
        })
        
        # 关闭数据集，释放内存
        print("    Closing datasets...")
        pr_data.close()
        prsn_data.close()
        snw_data.close()
        tas_data.close()
        
        # 强制垃圾回收
        gc.collect()
        
        return result_ds
        
    except Exception as e:
        print(f"    Error processing {model}/{experiment}: {str(e)}")
        # 尝试关闭可能已打开的数据集
        try:
            if 'pr_data' in locals():
                pr_data.close()
            if 'prsn_data' in locals():
                prsn_data.close()
            if 'snw_data' in locals():
                snw_data.close()
            if 'tas_data' in locals():
                tas_data.close()
        except:
            pass
        gc.collect()
        return None

# 处理每个模式
for model in tqdm(model_list):
    print(f"Processing model: {model}")
    
    # 处理历史数据
    print(f"Processing historical data for {model}")
    hist_ds = process_experiment(model, 'hist')
    
    if hist_ds is not None:
        # 保存历史数据
        hist_file = f"/hazards_single_models/hist/{model}_hist.nc"
        print(f"  Saving hist results to {hist_file}...")
        try:
            hist_ds.to_netcdf(hist_file)
            print(f"  Historical data for {model} saved successfully")
        except Exception as e:
            print(f"  Error saving historical data for {model}: {str(e)}")
        
        # 清理历史数据，释放内存
        hist_ds.close()
        del hist_ds
        gc.collect()
    else:
        print(f"  No valid historical data for {model}, skipping...")
        continue  # 如果没有历史数据，跳过此模式的其余处理
    
    # 处理未来数据 - ssp245
    future_ds = xr.Dataset()
    
    ssp245_ds = process_experiment(model, 'ssp245')
    if ssp245_ds is not None:
        future_ds = xr.merge([future_ds, ssp245_ds])
        ssp245_ds.close()
        del ssp245_ds
        gc.collect()
    
    # 处理未来数据 - ssp585
    ssp585_ds = process_experiment(model, 'ssp585')
    if ssp585_ds is not None:
        future_ds = xr.merge([future_ds, ssp585_ds])
        ssp585_ds.close()
        del ssp585_ds
        gc.collect()
    
    # 保存未来数据
    if not future_ds.data_vars:
        print(f"  No future scenario data for {model}")
    else:
        future_file = f"/hazards_single_models/projection/{model}_projection.nc"
        print(f"  Saving future results to {future_file}...")
        try:
            future_ds.to_netcdf(future_file)
            print(f"  Future data for {model} saved successfully")
        except Exception as e:
            print(f"  Error saving future data for {model}: {str(e)}")
        
        # 清理未来数据，释放内存
        future_ds.close()
        del future_ds
        gc.collect()
    
    print("="*50 + f" Completed {model} " + "="*50)

print("All models processed successfully.")
