未分類

#include <pybind11/pybind11.h>
#include <pybind11/numpy.h>
#include <vector>
#include <cmath>
#include <algorithm>

namespace py = pybind11;

// ==========================================
// 1. MEPSアンサンブル 確率・パーセンタイル一括計算 (シングルパス)
// ==========================================
py::dict calc_stats_v2(py::array_t<double> data) {
    py::buffer_info buf = data.request();
    if (buf.ndim != 3) throw std::runtime_error("Input must be a 3D array (members, lat, lon)");
    
    int num_members = buf.shape[0];
    int rows = buf.shape[1];
    int cols = buf.shape[2];
    int grid_size = rows * cols;
    double* ptr = static_cast<double*>(buf.ptr);
    
    auto p10 = py::array_t<double>({rows, cols});
    auto p25 = py::array_t<double>({rows, cols});
    auto p75 = py::array_t<double>({rows, cols});
    auto p90 = py::array_t<double>({rows, cols});
    auto prob1 = py::array_t<double>({rows, cols});
    auto prob10 = py::array_t<double>({rows, cols});
    auto prob20 = py::array_t<double>({rows, cols});
    auto prob30 = py::array_t<double>({rows, cols});
    
    double* ptr_p10 = static_cast<double*>(p10.request().ptr);
    double* ptr_p25 = static_cast<double*>(p25.request().ptr);
    double* ptr_p75 = static_cast<double*>(p75.request().ptr);
    double* ptr_p90 = static_cast<double*>(p90.request().ptr);
    double* ptr_prob1 = static_cast<double*>(prob1.request().ptr);
    double* ptr_prob10 = static_cast<double*>(prob10.request().ptr);
    double* ptr_prob20 = static_cast<double*>(prob20.request().ptr);
    double* ptr_prob30 = static_cast<double*>(prob30.request().ptr);
    
    int idx_p10 = std::max(0, (int)std::round((num_members - 1) * 0.10));
    int idx_p25 = std::max(0, (int)std::round((num_members - 1) * 0.25));
    int idx_p75 = std::min(num_members - 1, (int)std::round((num_members - 1) * 0.75));
    int idx_p90 = std::min(num_members - 1, (int)std::round((num_members - 1) * 0.90));

    #pragma omp parallel for
    for (int i = 0; i < grid_size; ++i) {
        std::vector<double> vals(num_members);
        int valid_count = 0;
        for (int m = 0; m < num_members; ++m) {
            double v = ptr[m * grid_size + i];
            vals[m] = v;
            if (!std::isnan(v)) valid_count++;
        }
        
        std::sort(vals.begin(), vals.end(), [](double a, double b) {
            if (std::isnan(a)) return false;
            if (std::isnan(b)) return true;
            return a < b;
        });
        
        ptr_p10[i] = vals[idx_p10];
        ptr_p25[i] = vals[idx_p25];
        ptr_p75[i] = vals[idx_p75];
        ptr_p90[i] = vals[idx_p90];
        
        int c1 = 0, c10 = 0, c20 = 0, c30 = 0;
        for (double v : vals) {
            if (!std::isnan(v)) {
                if (v >= 1.0) c1++;
                if (v >= 10.0) c10++;
                if (v >= 20.0) c20++;
                if (v >= 30.0) c30++;
            }
        }
        
        if (valid_count > 0) {
            ptr_prob1[i] = (double)c1 / valid_count * 100.0;
            ptr_prob10[i] = (double)c10 / valid_count * 100.0;
            ptr_prob20[i] = (double)c20 / valid_count * 100.0;
            ptr_prob30[i] = (double)c30 / valid_count * 100.0;
        } else {
            ptr_prob1[i] = std::nan(""); ptr_prob10[i] = std::nan("");
            ptr_prob20[i] = std::nan(""); ptr_prob30[i] = std::nan("");
        }
    }
    
    py::dict result;
    result["p10"] = p10; result["p25"] = p25; result["p75"] = p75; result["p90"] = p90;
    result["prob1"] = prob1; result["prob10"] = prob10; result["prob20"] = prob20; result["prob30"] = prob30;
    return result;
}

// ==========================================
// 2. 異種グリッド結合のインデックスマッピング (O(N)アルゴリズム)
// ==========================================
void map_onto_canvas(py::array_t<double> canvas, 
                     py::array_t<double> canvas_lon, py::array_t<double> canvas_lat,
                     py::array_t<double> src_val, 
                     py::array_t<double> src_lon, py::array_t<double> src_lat) {
    
    auto buf_canvas = canvas.request(); auto buf_clon = canvas_lon.request(); auto buf_clat = canvas_lat.request();
    auto buf_src = src_val.request(); auto buf_slon = src_lon.request(); auto buf_slat = src_lat.request();
    
    double* ptr_canvas = static_cast<double*>(buf_canvas.ptr);
    double* ptr_clon = static_cast<double*>(buf_clon.ptr); double* ptr_clat = static_cast<double*>(buf_clat.ptr);
    double* ptr_src = static_cast<double*>(buf_src.ptr);
    double* ptr_slon = static_cast<double*>(buf_slon.ptr); double* ptr_slat = static_cast<double*>(buf_slat.ptr);
    
    int c_lon_len = buf_clon.shape[0]; int c_lat_len = buf_clat.shape[0];
    int s_lon_len = buf_slon.shape[0]; int s_lat_len = buf_slat.shape[0];

    std::vector<int> lon_map(s_lon_len, -1);
    int c_idx = 0;
    for (int s_idx = 0; s_idx < s_lon_len; ++s_idx) {
        while (c_idx < c_lon_len && std::abs(ptr_clon[c_idx] - ptr_slon[s_idx]) > 1e-4) c_idx++;
        if (c_idx < c_lon_len) lon_map[s_idx] = c_idx;
        else break;
    }

    std::vector<int> lat_map(s_lat_len, -1);
    bool lat_descending = (c_lat_len > 1 && ptr_clat[0] > ptr_clat[c_lat_len - 1]);
    for (int s_idx = 0; s_idx < s_lat_len; ++s_idx) {
        double target = ptr_slat[s_idx];
        int left = 0, right = c_lat_len - 1;
        while (left <= right) {
            int mid = left + (right - left) / 2;
            if (std::abs(ptr_clat[mid] - target) < 1e-4) { lat_map[s_idx] = mid; break; }
            if (lat_descending) {
                if (ptr_clat[mid] > target) left = mid + 1; else right = mid - 1;
            } else {
                if (ptr_clat[mid] < target) left = mid + 1; else right = mid - 1;
            }
        }
    }

    for (int i = 0; i < s_lat_len; ++i) {
        int dst_lat_idx = lat_map[i];
        if (dst_lat_idx == -1) continue;
        for (int j = 0; j < s_lon_len; ++j) {
            int dst_lon_idx = lon_map[j];
            if (dst_lon_idx == -1) continue;
            double val = ptr_src[i * s_lon_len + j];
            if (!std::isnan(val)) ptr_canvas[dst_lat_idx * c_lon_len + dst_lon_idx] = val;
        }
    }
}

// ==========================================
// 3. 熱力学パラメーター計算 (ループフュージョン)
// ==========================================
py::tuple calc_thermo(py::array_t<double> tc_arr, py::array_t<double> rh_arr, double level) {
    auto buf_tc = tc_arr.request(); auto buf_rh = rh_arr.request();
    int size = buf_tc.size;
    double* ptr_tc = static_cast<double*>(buf_tc.ptr);
    double* ptr_rh = static_cast<double*>(buf_rh.ptr);
    
    auto tddep_arr = py::array_t<double>(buf_tc.shape);
    auto ep_arr = py::array_t<double>(buf_tc.shape);
    double* ptr_tddep = static_cast<double*>(tddep_arr.request().ptr);
    double* ptr_ep = static_cast<double*>(ep_arr.request().ptr);
    
    #pragma omp parallel for
    for (int i = 0; i < size; ++i) {
        double tc = ptr_tc[i]; double rh = ptr_rh[i];
        if (std::isnan(tc) || std::isnan(rh)) {
            ptr_tddep[i] = std::nan(""); ptr_ep[i] = std::nan(""); continue;
        }
        double rh_c = std::max(0.1, std::min(100.0, rh));
        double e = 6.112 * std::exp((17.67 * tc) / (tc + 243.5)) * (rh_c / 100.0);
        double log_e = std::log(e / 6.112);
        double td = (243.5 * log_e) / (17.67 - log_e);
        
        ptr_tddep[i] = tc - td;
        
        double tk = tc + 273.15;
        double theta = tk * std::pow(1000.0 / level, 0.2854);
        double w = 0.622 * e / (level - e);
        ptr_ep[i] = theta * std::exp((2.5e6 * w) / (1004.0 * tk));
    }
    return py::make_tuple(tddep_arr, ep_arr);
}

// 既存の渦度計算
py::array_t<double> calc_vorticity(py::array_t<double> u_arr, py::array_t<double> v_arr, py::array_t<double> lon_arr, py::array_t<double> lat_arr) {
    auto buf_u = u_arr.request(); auto buf_v = v_arr.request();
    auto buf_lon = lon_arr.request(); auto buf_lat = lat_arr.request();
    int rows = buf_u.shape[0]; int cols = buf_u.shape[1];
    auto vort_arr = py::array_t<double>({rows, cols});
    double* ptr_vort = static_cast<double*>(vort_arr.request().ptr);
    double* u = static_cast<double*>(buf_u.ptr); double* v = static_cast<double*>(buf_v.ptr);
    double* lon = static_cast<double*>(buf_lon.ptr); double* lat = static_cast<double*>(buf_lat.ptr);

    const double R = 6371000.0; const double PI = 3.14159265358979323846;

    for (int i = 0; i < rows; ++i) {
        for (int j = 0; j < cols; ++j) {
            if (i == 0 || i == rows - 1 || j == 0 || j == cols - 1) { ptr_vort[i * cols + j] = 0.0; continue; }
            double dlat = (lat[i + 1] - lat[i - 1]) * PI / 180.0;
            double dlon = (lon[j + 1] - lon[j - 1]) * PI / 180.0;
            double lat_rad = lat[i] * PI / 180.0;
            double dy = R * dlat; double dx = R * std::cos(lat_rad) * dlon;
            if (dx == 0) dx = 1e-10; if (dy == 0) dy = 1e-10;
            double dv_dx = (v[i * cols + j + 1] - v[i * cols + j - 1]) / dx;
            double du_dy = (u[(i + 1) * cols + j] - u[(i - 1) * cols + j]) / dy;
            ptr_vort[i * cols + j] = (dv_dx - du_dy) * 1e5;
        }
    }
    return vort_arr;
}

PYBIND11_MODULE(fast_meteo, m) {
    m.def("calc_stats", &calc_stats_v2, "Calculate MEPS statistics");
    m.def("map_onto_canvas", &map_onto_canvas, "Fast grid merging");
    m.def("calc_thermo", &calc_thermo, "Calculate td, tddep, ep");
    m.def("calc_vorticity", &calc_vorticity, "Calculate vorticity");
}