Team Ai
Apppublic

nakas/waveVisualizer-Codex

sourceHugging Facecc-by-4.0updated 1y agoView on Hugging Face
0likes
grib_wave_puller.py303 linesDownload Raw Back to backend
1"""2Vendored GRIBWavePuller from NWPS_SWAN project (trimmed only for runtime import here).3If Arctic-specific helpers are missing, the puller will log and continue with fallbacks.4"""5 6import os7import sys8import tempfile9import logging10import subprocess11from datetime import datetime, timedelta12import numpy as np13import xarray as xr14from ecmwf.opendata import Client15import requests16 17logging.basicConfig(level=logging.INFO)18logger = logging.getLogger(__name__)19 20# Optional Arctic handler21try:22    from arctic_grib_handler import ArcticGRIBHandler  # type: ignore23    ARCTIC_HANDLER_AVAILABLE = True24    logger.info("✅ Arctic GRIB Handler loaded (Docker production version)")25except Exception as e:26    logger.warning(f"Arctic GRIB handler not available: {e}")27    ARCTIC_HANDLER_AVAILABLE = False28 29 30class GRIBWavePuller:31    def __init__(self):32        self.client = Client("ecmwf")33        self.output_dir = os.getenv('OUTPUT_DIR', '/tmp/wave_data')34        os.makedirs(self.output_dir, exist_ok=True)35        self._setup_eccodes_environment()36 37    def _setup_eccodes_environment(self):38        try:39            os.environ['ECCODES_GRIB_STRICT_PARSING'] = '0'40            os.environ['ECCODES_GRIB_IGNORE_GRID_DEFINITION'] = '1'41        except Exception:42            pass43 44    def fetch_ecmwf_wave_grib(self, forecast_time=0):45        try:46            temp_file = tempfile.NamedTemporaryFile(delete=False, suffix='.grib2')47            try:48                self.client.retrieve(49                    type="fc",50                    param=["swh"],51                    time=0,52                    step=forecast_time,53                    target=temp_file.name,54                )55                return temp_file.name56            except Exception:57                if os.path.exists(temp_file.name):58                    os.unlink(temp_file.name)59                return None60        except Exception:61            return None62 63    def fetch_noaa_wave_grib(self, forecast_hour=0):64        """Fetch NOAA WW3 files (regional + global attempt). Returns list of (path, region, run, fh)."""65        try:66            base_url = "https://nomads.ncep.noaa.gov/pub/data/nccf/com/gfs/prod"67            now = datetime.utcnow()68            dates_to_try = [69                now.strftime("%Y%m%d"),70                (now - timedelta(days=1)).strftime("%Y%m%d"),71                (now - timedelta(days=2)).strftime("%Y%m%d"),72            ]73            if now.hour >= 18:74                preferred_runs = ["18", "12", "06", "00"]75            elif now.hour >= 12:76                preferred_runs = ["12", "06", "00", "18"]77            elif now.hour >= 6:78                preferred_runs = ["06", "00", "18", "12"]79            else:80                preferred_runs = ["00", "18", "12", "06"]81 82            for date_str in dates_to_try:83                for hour in preferred_runs:84                    successful = []85                    regional_files = [86                        (f"gfswave.t{hour}z.atlocn.0p16.f{forecast_hour:03d}.grib2", "Atlantic"),87                        (f"gfswave.t{hour}z.epacif.0p16.f{forecast_hour:03d}.grib2", "East_Pacific"),88                        (f"gfswave.t{hour}z.arctic.9km.f{forecast_hour:03d}.grib2", "Arctic"),89                        (f"gfswave.t{hour}z.global.0p16.f{forecast_hour:03d}.grib2", "Global"),90                    ]91                    for filename, region_name in regional_files:92                        url = f"{base_url}/gfs.{date_str}/{hour}/wave/gridded/{filename}"93                        try:94                            tf = tempfile.NamedTemporaryFile(delete=False, suffix='.grib2')95                            r = requests.get(url, timeout=300)96                            if r.status_code == 200:97                                tf.write(r.content)98                                tf.close()99                                successful.append((tf.name, region_name, hour, forecast_hour))100                            else:101                                tf.close(); os.unlink(tf.name)102                        except Exception:103                            try:104                                tf.close(); os.unlink(tf.name)105                            except Exception:106                                pass107                            continue108                    if successful:109                        return successful110            return None111        except Exception:112            return None113 114    def process_grib_file(self, grib_file_path, region_name=None):115        try:116            ds = xr.open_dataset(grib_file_path, engine='cfgrib', decode_timedelta=True)117        except Exception:118            return None, None119 120        # Identify variables121        vars_map = {name: ds[name] for name in ds.variables}122        wave_var = None123        for cand in ['swh', 'HTSGW', 'htsgw']:124            if cand in vars_map:125                wave_var = cand; break126        if wave_var is None:127            # fallback heuristic128            for n in vars_map:129                if 'wave' in n.lower() and 'height' in n.lower():130                    wave_var = n; break131        if wave_var is None:132            ds.close()133            return None, None134 135        wave_heights = vars_map[wave_var].values136 137        wave_dir = None138        for cand in ['dirpw', 'DIRPW', 'dp', 'wvdir', 'WVDIR', 'dir', 'mwd', 'MWD', 'MWDIR']:139            if cand in vars_map:140                wave_dir = vars_map[cand].values; break141 142        wave_per = None143        for cand in ['perpw', 'PERPW', 'tp', 'wvper', 'WVPER', 'per', 'pp1d', 'PP1D', 'mwp', 'MWP']:144            if cand in vars_map:145                wave_per = vars_map[cand].values; break146 147        lats = ds.latitude.values if 'latitude' in ds else ds.lat.values148        lons = ds.longitude.values if 'longitude' in ds else ds.lon.values149 150        # Sample points (downsample for visualization)151        lon_grid, lat_grid = np.meshgrid(lons, lats)152        flat_lats = lat_grid.flatten()153        flat_lons = lon_grid.flatten()154        flat_waves = wave_heights.flatten()155        mask = ~np.isnan(flat_waves)156        if wave_dir is not None:157            mask &= ~np.isnan(wave_dir.flatten())158        idx = np.random.choice(np.where(mask)[0], size=min(1000, mask.sum()), replace=False) if mask.any() else np.array([])159 160        points = []161        for i in idx:162            point = {163                'lat': float(flat_lats[i]),164                'lon': float(flat_lons[i]),165                'wave_height': float(flat_waves[i]),166            }167            if wave_dir is not None:168                d = float(wave_dir.flatten()[i])169                point['wave_direction'] = d170                mag = point['wave_height'] * 0.1171                rad = np.deg2rad(d)172                point['u_component'] = float(mag * np.sin(rad))173                point['v_component'] = float(-mag * np.cos(rad))174            if wave_per is not None:175                point['wave_period'] = float(wave_per.flatten()[i])176            points.append(point)177 178        data = {179            'timestamp': datetime.utcnow().isoformat(),180            'data_source': 'NOAA_GRIB' if 'gfswave' in os.path.basename(grib_file_path) else 'GRIB',181            'parameters_found': {182                'wave_height': wave_var,183                'wave_direction': 'present' if wave_dir is not None else None,184                'wave_period': 'present' if wave_per is not None else None,185                'has_velocity_components': wave_dir is not None,186            },187            'grid_info': {188                'lat_min': float(np.nanmin(lats)),189                'lat_max': float(np.nanmax(lats)),190                'lon_min': float(np.nanmin(lons)),191                'lon_max': float(np.nanmax(lons)),192                'grid_shape': list(wave_heights.shape),193            },194            'sample_points': points,195        }196        # Optional: include a downsampled U/V grid for velocity layers197        try:198            if wave_dir is not None:199                # Compute U/V on the native grid200                # Determine reasonable downsample strides to keep <= ~360x180201                ny, nx = wave_heights.shape202                sy = max(1, ny // 180)203                sx = max(1, nx // 360)204                lats_ds = lats[::sy]205                lons_ds = lons[::sx]206                # Align 2D arrays for downsample207                wh_ds = wave_heights[::sy, ::sx]208                wd_ds = wave_dir[::sy, ::sx]209                wp_ds = None210                if wave_per is not None:211                    try:212                        wp_ds = wave_per[::sy, ::sx]213                    except Exception:214                        wp_ds = None215                # Compute U/V216                dir_rad = np.deg2rad(wd_ds)217                # Base speed from period if present (deep water group velocity)218                if wp_ds is not None:219                    base = 0.78 * np.clip(wp_ds, 0, 20)220                else:221                    base = 1.0 + 0.2 * np.clip(wh_ds, 0, 10)222 223                # Add spatial variation via normalized wave height224                try:225                    p50 = float(np.nanpercentile(wh_ds, 50))226                    p90 = float(np.nanpercentile(wh_ds, 90))227                    denom = (p90 - p50) if (p90 - p50) > 1e-6 else 1.0228                    hnorm = np.clip((wh_ds - p50) / denom, -1.0, 2.0)229                except Exception:230                    hnorm = 0.0231                mag = base * (1.0 + 0.5 * hnorm)232 233                # Clamp to a reasonable range for visualization234                mag = np.clip(mag, 0.0, 15.0)235                u_ds = mag * np.sin(dir_rad)236                v_ds = -mag * np.cos(dir_rad)237                data['grid_uv'] = {238                    'lats': lats_ds.tolist() if hasattr(lats_ds, 'tolist') else list(map(float, lats_ds)),239                    'lons': lons_ds.tolist() if hasattr(lons_ds, 'tolist') else list(map(float, lons_ds)),240                    'u': np.nan_to_num(u_ds, nan=0.0, posinf=0.0, neginf=0.0).tolist(),241                    'v': np.nan_to_num(v_ds, nan=0.0, posinf=0.0, neginf=0.0).tolist(),242                }243                try:244                    sp = np.sqrt(u_ds*u_ds + v_ds*v_ds)245                    data['grid_uv_info'] = {246                        'speed_min': float(np.nanmin(sp)),247                        'speed_max': float(np.nanmax(sp)),248                        'speed_mean': float(np.nanmean(sp)),249                    }250                except Exception:251                    pass252        except Exception:253            # If any step fails, just skip embedding grid_uv254            pass255        ds.close()256        return data, grib_file_path257 258    def process_multiple_regional_files(self, regional_files):259        combined = []260        for path, region_name, *_ in regional_files:261            try:262                res, _ = self.process_grib_file(path, region_name=region_name)263                if res and 'sample_points' in res:264                    combined.extend(res['sample_points'])265            finally:266                try:267                    if os.path.exists(path):268                        os.unlink(path)269                except Exception:270                    pass271        if not combined:272            return None273        return {274            'timestamp': datetime.utcnow().isoformat(),275            'data_source': 'NOAA_MULTI_REGIONAL_GRIB',276            'parameters_found': {'has_velocity_components': True},277            'grid_info': {},278            'sample_points': combined,279        }280 281    def fetch_global_wave_data(self, forecast_hour=0):282        result = self.fetch_noaa_wave_grib(forecast_hour)283        if isinstance(result, list) and result:284            if any(region == 'Global' for _, region, *_ in result):285                # Prefer the global grid if present286                global_entry = next((t for t in result if t[1] == 'Global'), None)287                if global_entry:288                    data, _ = self.process_grib_file(global_entry[0], region_name='Global')289                    return data290            # Otherwise combine sample points from regions291            return self.process_multiple_regional_files(result)292 293        # Fallback ECMWF (may not include waves)294        grib_file = self.fetch_ecmwf_wave_grib(forecast_hour)295        if grib_file:296            data, _ = self.process_grib_file(grib_file)297            try:298                os.unlink(grib_file)299            except Exception:300                pass301            return data302        return None303