diff --git a/covjsonkit/encoder/encoder.py b/covjsonkit/encoder/encoder.py index 0332de2..a266bfb 100644 --- a/covjsonkit/encoder/encoder.py +++ b/covjsonkit/encoder/encoder.py @@ -12,6 +12,69 @@ from covjsonkit.param_db import get_param_ids, get_params, get_units +try: + # Polytope compacts unstructured-grid (e.g. ICON, Lambert LAM) leaves into a single + # MergedTensorIndexNode holding axes=(lat_axis, lon_axis) and values=(lat, lon). + from polytope_feature.datacube.tensor_index_tree import MergedTensorIndexNode +except ImportError: # older polytope without merged nodes + MergedTensorIndexNode = None + +try: + # Polytope compacts all the lat/lon points under one path into a single + # BulkMergedTensorIndexNode (unstructured grids), or its BulkGridTensorIndexNode + # subclass (structured grids), holding ``coordinates`` (N, 2) and one result array + # of N values per combination of the compressed axes above it. + from polytope_feature.datacube.tensor_index_tree import BulkMergedTensorIndexNode +except ImportError: # older polytope without bulk nodes + BulkMergedTensorIndexNode = None + + +def is_bulk_node(node) -> bool: + """True if ``node`` is a polytope ``BulkMergedTensorIndexNode`` (an array-backed lat/lon leaf). + + Falls back to duck-typing (``.coordinates`` and ``.point_count`` present) if the + polytope import was unavailable. + """ + if BulkMergedTensorIndexNode is not None: + return isinstance(node, BulkMergedTensorIndexNode) + return hasattr(node, "coordinates") and hasattr(node, "point_count") + + +def is_merged_node(node) -> bool: + """True if ``node`` is a polytope ``MergedTensorIndexNode`` (a compacted lat/lon leaf). + + Such nodes carry ``axes=(lat_axis, lon_axis)`` and ``values=(lat, lon)`` for a single + spatial point, and are always leaves. Bulk nodes, although a subclass, are not + single points and are excluded. Falls back to duck-typing (``.axes`` present) + if the polytope import was unavailable. + """ + if is_bulk_node(node): + return False + if MergedTensorIndexNode is not None: + return isinstance(node, MergedTensorIndexNode) + return hasattr(node, "axes") and getattr(node, "axes", None) is not None + + +def bulk_flat_result(node) -> list: + """Flatten a bulk node's result into the layout of a legacy leaf holding all its points. + + The legacy layout is combination-major: one block per combination of the compressed + axes, each block holding the values of all the points. + """ + if len(node.result) == 0: + return [None] * node.point_count + return np.concatenate([np.asarray(values, dtype=object).reshape(-1) for values in node.result]).tolist() + + +def bulk_point_results(node): + """Yield ``(lat, lon, result)`` per point of a bulk node, as if it were a single-point leaf.""" + if len(node.result) == 0: + values = np.full((1, node.point_count), None, dtype=object) + else: + values = np.stack([np.asarray(r, dtype=object).reshape(-1) for r in node.result]) + for i, (lat, lon) in enumerate(node.coordinates.tolist()): + yield lat, lon, values[:, i].tolist() + def timedelta_to_step_string(td: timedelta) -> str: """ @@ -378,37 +441,21 @@ def calculate_index_bounds(level_len, num_len, para_len, step_len, l, i, j, k): end_index = start_index + int(step_len) return start_index, end_index - def append_composite_coords(dates, tree_values, lat, coords): + def append_composite_coords(dates, points, coords): # for date in dates: - for value in tree_values: - coords[dates]["composite"].append([lat, value]) - - if len(tree.children) != 0: - for child in tree.children: - handle_non_leaf_node(child) - result = handle_specific_axes(child) - if result is not None: - if child.axis.name == "latitude": - fields["lat"] = result - elif child.axis.name == "levelist": - fields["levels"] = result - if "l" in fields: - fields["l"].extend(result) - elif child.axis.name == "param": - fields["param"] = result - elif child.axis.name in [date_key, "time"]: - fields["dates"].extend(result) - elif child.axis.name == "number": - fields["number"] = result - elif child.axis.name == "step": - fields["step"] = result - if "s" in fields: - fields["s"].extend(result) - - self.walk_tree(child, fields, coords, mars_metadata, range_dict, date_key=date_key) - else: - tree.values = [float(val) for val in tree.values] - if all(val is None for val in tree.result): + for lat, lon in points: + coords[dates]["composite"].append([lat, lon]) + + def emit_leaf(points, result): + """Emit one spatial leaf: append [lat, lon] composite coords and slice results. + + Shared by the legacy longitude-leaf path (``points`` pairs the parent latitude + with each of the leaf's longitudes), the compacted ``MergedTensorIndexNode`` + path (a single point) and the ``BulkMergedTensorIndexNode`` path (all its + points, with ``result`` flattened by :func:`bulk_flat_result`). + """ + points = [(lat, float(lon)) for lat, lon in points] + if all(val is None for val in result): fields["dates"] = fields["dates"][:-1] for date in fields["dates"]: for level in fields["levels"]: @@ -419,13 +466,13 @@ def append_composite_coords(dates, tree_values, lat, coords): if key in range_dict: del range_dict[key] else: - tree.result = [float(val) if val is not None else val for val in tree.result] - level_len = len(tree.result) / len(fields["levels"]) + result = [float(val) if val is not None else val for val in result] + level_len = len(result) / len(fields["levels"]) num_len = level_len / len(fields["number"]) para_len = num_len / len(fields["param"]) step_len = para_len / len(fields["step"]) - append_composite_coords(fields["dates"][-1], tree.values, fields["lat"], coords) + append_composite_coords(fields["dates"][-1], points, coords) for l, level in enumerate(fields["levels"]): # noqa: E741 for i, num in enumerate(fields["number"]): @@ -437,7 +484,41 @@ def append_composite_coords(dates, tree_values, lat, coords): key = create_composite_key(fields["dates"][-1], level, num, para, s) if key not in range_dict: range_dict[key] = [] - range_dict[key].extend(tree.result[start_index:end_index]) + range_dict[key].extend(result[start_index:end_index]) + + if len(tree.children) != 0: + for child in tree.children: + # Bulk leaf: all points under this path at once. + if is_bulk_node(child): + emit_leaf(child.coordinates.tolist(), bulk_flat_result(child)) + continue + # Compacted unstructured leaf: values=(lat, lon), own result. Emit directly. + if is_merged_node(child): + emit_leaf([(child.values[0], child.values[1])], child.result) + continue + handle_non_leaf_node(child) + result = handle_specific_axes(child) + if result is not None: + if child.axis.name == "latitude": + fields["lat"] = result + elif child.axis.name == "levelist": + fields["levels"] = result + if "l" in fields: + fields["l"].extend(result) + elif child.axis.name == "param": + fields["param"] = result + elif child.axis.name in [date_key, "time"]: + fields["dates"].extend(result) + elif child.axis.name == "number": + fields["number"] = result + elif child.axis.name == "step": + fields["step"] = result + if "s" in fields: + fields["s"].extend(result) + + self.walk_tree(child, fields, coords, mars_metadata, range_dict, date_key=date_key) + else: + emit_leaf([(fields["lat"], lon) for lon in tree.values], tree.result) def walk_tree_step(self, tree, fields, coords, mars_metadata, range_dict): def create_composite_key_step(date, level, num, para): @@ -496,34 +577,10 @@ def append_composite_coords_step(dates, tree_values, lat, coords): for value in tree_values: coords[dates]["composite"].append([lat, value]) - if len(tree.children) != 0: - for child in tree.children: - handle_non_leaf_node_step(child) - result = handle_specific_axes_step(child) - if result is not None: - if child.axis.name == "latitude": - fields["lat"] = result - elif child.axis.name == "levelist": - fields["levels"] = result - if "l" in fields: - fields["l"].extend(result) - elif child.axis.name == "param": - fields["param"] = result - elif child.axis.name in ["date"]: - fields["dates"].extend(result) - elif child.axis.name == "number": - fields["number"] = result - elif child.axis.name == "step": - fields["step"] = result - if "s" in fields: - fields["s"].extend(result) - elif child.axis.name == "time": - fields["times"].extend(result) - - self.walk_tree_step(child, fields, coords, mars_metadata, range_dict) - else: - tree.values = [float(val) for val in tree.values] - if all(val is None for val in tree.result): + def emit_leaf_step(lat, lon_values, result): + """Emit one spatial leaf for the step walker (shared by legacy and merged).""" + lon_values = [float(val) for val in lon_values] + if all(val is None for val in result): fields["dates"] = fields["dates"][:-1] for date in fields["dates"]: for level in fields["levels"]: @@ -534,61 +591,66 @@ def append_composite_coords_step(dates, tree_values, lat, coords): if key in range_dict: del range_dict[key] else: - tree.result = [float(val) if val is not None else val for val in tree.result] - date_len = len(tree.result) / len(fields["dates"]) + result = [float(val) if val is not None else val for val in result] + date_len = len(result) / len(fields["dates"]) level_len = date_len / len(fields["levels"]) para_len = level_len / len(fields["param"]) - # time_len = para_len / len(fields["times"]) - # coords_len = len(tree.values) for date in fields["dates"]: - append_composite_coords_step(date, tree.values, fields["lat"], coords) - """ - for ti, _ in enumerate(fields["times"]): - for d, date in enumerate(fields["dates"]): - for l, level in enumerate(fields["levels"]): # noqa: E741 - for i, num in enumerate(fields["number"]): - for j, para in enumerate(fields["param"]): - # for k, t in enumerate(fields["times"]): - # start_index, end_index = calculate_index_bounds_step( - # level_len, num_len, para_len, time_len, l, i, j, k - # ) - key = create_composite_key_step(date, level, num, para) - if key not in range_dict: - range_dict[key] = [] - # range_dict[key].extend(tree.result[start_index:end_index]) - # print(d, date_len,j, para_len) - # print(d*date_len+j*para_len) - # print(int(d*date_len+j*para_len+len(fields["times"]))) - # print(tree.result[int(d*date_len+j*para_len+len)]) - #print(tree.result) - range_dict[key].append(#tree.result - tree.result[ - int(d * date_len + j * para_len + ti * time_len) : int( - d * date_len + j * para_len + ti * time_len + len(tree.values) - ) - ] - ) - """ + append_composite_coords_step(date, lon_values, lat, coords) + for d, date in enumerate(fields["dates"]): for l, level in enumerate(fields["levels"]): # noqa: E741 for i, num in enumerate(fields["number"]): for j, para in enumerate(fields["param"]): - # for k, t in enumerate(fields["times"]): - # start_index, end_index = calculate_index_bounds_step( - # level_len, num_len, para_len, time_len, l, i, j, k - # ) key = create_composite_key_step(date, level, num, para) if key not in range_dict: range_dict[key] = [] - range_dict[key].append( # tree.result - tree.result[ + range_dict[key].append( + result[ int(d * date_len + l * level_len + j * para_len) : int( d * date_len + l * level_len + j * para_len + len(fields["times"]) ) ] ) + if len(tree.children) != 0: + for child in tree.children: + # Bulk leaf: emit each point as a compacted single-point leaf. + if is_bulk_node(child): + for lat, lon, result in bulk_point_results(child): + emit_leaf_step(lat, [lon], result) + continue + # Compacted unstructured leaf: values=(lat, lon), own result. Emit directly. + if is_merged_node(child): + emit_leaf_step(child.values[0], [child.values[1]], child.result) + continue + handle_non_leaf_node_step(child) + result = handle_specific_axes_step(child) + if result is not None: + if child.axis.name == "latitude": + fields["lat"] = result + elif child.axis.name == "levelist": + fields["levels"] = result + if "l" in fields: + fields["l"].extend(result) + elif child.axis.name == "param": + fields["param"] = result + elif child.axis.name in ["date"]: + fields["dates"].extend(result) + elif child.axis.name == "number": + fields["number"] = result + elif child.axis.name == "step": + fields["step"] = result + if "s" in fields: + fields["s"].extend(result) + elif child.axis.name == "time": + fields["times"].extend(result) + + self.walk_tree_step(child, fields, coords, mars_metadata, range_dict) + else: + emit_leaf_step(fields["lat"], tree.values, tree.result) + def walk_tree_month(self, tree, fields, coords, mars_metadata, range_dict, _ctx=None): """Walk the result tree for monthly-mean streams (e.g. clmn). @@ -640,47 +702,13 @@ def append_composite_coords_month(date_key, tree_values, lat): for value in tree_values: coords[date_key]["composite"].append([lat, value]) - if len(tree.children) != 0: - for child in tree.children: - handle_non_leaf_node_month(child) - result = handle_specific_axes_month(child) - - # Build a child context that inherits the current year/month - child_ctx = dict(_ctx) + def emit_leaf_month(lat, lon_values, result): + """Emit one spatial leaf for the month walker (shared by legacy and merged). - if result is not None: - if child.axis.name == "latitude": - fields["lat"] = result - elif child.axis.name == "levelist": - fields["levels"] = result - if "l" in fields: - fields["l"].extend(result) - elif child.axis.name == "param": - fields["param"] = result - elif child.axis.name == "year": - fields["years"] = list(result) if fields.get("years") == [] else fields["years"] - child_ctx["years"] = result - child_ctx["_axis_order"] = _ctx.get("_axis_order", []) + ["year"] - # If month is already fixed in context, register dates now. - if "months" in _ctx: - for y in result: - for m in _ctx["months"]: - _register_date_key(_year_month_key(y, m)) - elif child.axis.name == "month": - fields["months"] = list(result) if fields.get("months") == [] else fields["months"] - child_ctx["months"] = result - # If year is already fixed in context, register dates now. - if "years" in _ctx: - for y in _ctx["years"]: - for m in result: - _register_date_key(_year_month_key(y, m)) - # Track that month is the inner axis relative to year - child_ctx["_axis_order"] = _ctx.get("_axis_order", []) + ["month"] - elif child.axis.name == "number": - fields["number"] = result - - self.walk_tree_month(child, fields, coords, mars_metadata, range_dict, _ctx=child_ctx) - else: + ``lat`` is the latitude for this leaf, ``lon_values`` the leaf's list of + longitudes (length 1 for a compacted ``MergedTensorIndexNode``), and + ``result`` its flat value array. + """ # Leaf node — ensure all (year, month) combinations from context are registered. ctx_years = _ctx.get("years", fields.get("years", [])) ctx_months = _ctx.get("months", fields.get("months", [])) @@ -708,8 +736,8 @@ def append_composite_coords_month(date_key, tree_values, lat): else: leaf_dates = fields["dates"] - tree.values = [float(val) for val in tree.values] - if all(val is None for val in tree.result): + lon_values = [float(val) for val in lon_values] + if all(val is None for val in result): # Remove date entries for this leaf that produced no data. for key in leaf_dates: if key in fields["dates"]: @@ -721,20 +749,20 @@ def append_composite_coords_month(date_key, tree_values, lat): if rkey in range_dict: del range_dict[rkey] else: - tree.result = [float(val) if val is not None else val for val in tree.result] + result = [float(val) if val is not None else val for val in result] n_dates = len(leaf_dates) n_levels = len(fields["levels"]) n_params = len(fields["param"]) - date_len = len(tree.result) / n_dates if n_dates else len(tree.result) + date_len = len(result) / n_dates if n_dates else len(result) level_len = date_len / n_levels if n_levels else date_len para_len = level_len / n_params if n_params else level_len # Append this leaf's longitude values to composite coords for # every date key in scope. for date in leaf_dates: - append_composite_coords_month(date, tree.values, fields["lat"]) + append_composite_coords_month(date, lon_values, lat) for d, date in enumerate(leaf_dates): for l, level in enumerate(fields["levels"]): # noqa: E741 @@ -744,8 +772,60 @@ def append_composite_coords_month(date_key, tree_values, lat): if key not in range_dict: range_dict[key] = [] start = int(d * date_len + l * level_len + j * para_len) - end = int(start + len(tree.values)) - range_dict[key].append(tree.result[start:end]) + end = int(start + len(lon_values)) + range_dict[key].append(result[start:end]) + + if len(tree.children) != 0: + for child in tree.children: + # Bulk leaf: emit each point as a compacted single-point leaf. + if is_bulk_node(child): + for lat, lon, result in bulk_point_results(child): + emit_leaf_month(lat, [lon], result) + continue + # Compacted unstructured leaf: values=(lat, lon), own result. Emit directly. + if is_merged_node(child): + emit_leaf_month(child.values[0], [child.values[1]], child.result) + continue + handle_non_leaf_node_month(child) + result = handle_specific_axes_month(child) + + # Build a child context that inherits the current year/month + child_ctx = dict(_ctx) + + if result is not None: + if child.axis.name == "latitude": + fields["lat"] = result + elif child.axis.name == "levelist": + fields["levels"] = result + if "l" in fields: + fields["l"].extend(result) + elif child.axis.name == "param": + fields["param"] = result + elif child.axis.name == "year": + fields["years"] = list(result) if fields.get("years") == [] else fields["years"] + child_ctx["years"] = result + child_ctx["_axis_order"] = _ctx.get("_axis_order", []) + ["year"] + # If month is already fixed in context, register dates now. + if "months" in _ctx: + for y in result: + for m in _ctx["months"]: + _register_date_key(_year_month_key(y, m)) + elif child.axis.name == "month": + fields["months"] = list(result) if fields.get("months") == [] else fields["months"] + child_ctx["months"] = result + # If year is already fixed in context, register dates now. + if "years" in _ctx: + for y in _ctx["years"]: + for m in result: + _register_date_key(_year_month_key(y, m)) + # Track that month is the inner axis relative to year + child_ctx["_axis_order"] = _ctx.get("_axis_order", []) + ["month"] + elif child.axis.name == "number": + fields["number"] = result + + self.walk_tree_month(child, fields, coords, mars_metadata, range_dict, _ctx=child_ctx) + else: + emit_leaf_month(fields["lat"], tree.values, tree.result) @abstractmethod def add_coverage(self, mars_metadata, coords, values): diff --git a/tests/conftest.py b/tests/conftest.py index ce742e8..aff613a 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -2,6 +2,15 @@ from polytope_feature.datacube.datacube_axis import IntDatacubeAxis from polytope_feature.datacube.tensor_index_tree import TensorIndexTree +try: + # Only available on polytope versions that support compacted unstructured + # (ICON, Lambert LAM) results. Absent on released polytope, in which case + # the merged-node fixtures/tests below are skipped rather than breaking + # collection of the whole test suite. + from polytope_feature.datacube.tensor_index_tree import MergedTensorIndexNode +except ImportError: + MergedTensorIndexNode = None + # -- Shared constants for reforecast tests -- REFORECAST_METADATA_BASE = { @@ -61,14 +70,36 @@ def make_point(lat, lon, result): return lat_n -def forecast_tree(points, param="167", step=(0,), date=np.datetime64("2025-01-01T00:00:00")): +def make_merged_point(lat, lon, result): + """Create a compacted MergedTensorIndexNode for a single spatial point. + + This is the unstructured-grid equivalent of :func:`make_point`: instead of a + latitude node with a longitude-leaf child, the (lat, lon) pair is compacted + into a single leaf node carrying its own ``result``. + """ + if MergedTensorIndexNode is None: + raise RuntimeError( + "MergedTensorIndexNode is not available in this polytope build; " "merged-node tests should be skipped." + ) + lat_ax = IntDatacubeAxis() + lat_ax.name = "latitude" + lon_ax = IntDatacubeAxis() + lon_ax.name = "longitude" + merged = MergedTensorIndexNode(axes=(lat_ax, lon_ax), values=(lat, lon)) + merged.result = [np.float64(r) for r in result] + return merged + + +def forecast_tree(points, param="167", step=(0,), date=np.datetime64("2025-01-01T00:00:00"), point_factory=make_point): """Build a standard forecast TensorIndexTree with the given spatial points. Args: - points: list of (lat, lon, result_list) tuples, passed to make_point(). + points: list of (lat, lon, result_list) tuples, passed to point_factory(). param: MARS parameter code. step: tuple of step values. date: forecast date. + point_factory: callable(lat, lon, result) -> node. Use make_merged_point + to build a compacted unstructured (MergedTensorIndexNode) tree. """ tree = chain( TensorIndexTree(), @@ -84,11 +115,36 @@ def forecast_tree(points, param="167", step=(0,), date=np.datetime64("2025-01-01 ) parent = tip(tree) for lat, lon, result in points: - parent.add_child(make_point(lat, lon, result)) + parent.add_child(point_factory(lat, lon, result)) + return tree + + +def month_tree(points, param="167", years=(2020, 2021), months=(1,), point_factory=make_point): + """Build a monthly-mean tree (year/month axes) for the walk_tree_month path. + + Args: + points: list of (lat, lon, result_list) tuples, passed to point_factory(). + param: MARS parameter code. + years: tuple of year values. + months: tuple of month values. + point_factory: callable(lat, lon, result) -> node. Use make_merged_point + to build a compacted unstructured (MergedTensorIndexNode) tree. + """ + tree = chain( + TensorIndexTree(), + node("class", ("od",)), + node("levtype", ("sfc",)), + node("param", (param,)), + node("year", years), + node("month", months), + ) + parent = tip(tree) + for lat, lon, result in points: + parent.add_child(point_factory(lat, lon, result)) return tree -def reforecast_branch(hdate, points, param="167", step=(0,)): +def reforecast_branch(hdate, points, param="167", step=(0,), point_factory=make_point): """Build a reforecast branch rooted at an hdate node. Attaches spatial points at the leaf. Caller is responsible for @@ -106,7 +162,7 @@ def reforecast_branch(hdate, points, param="167", step=(0,)): ) parent = tip(branch) for lat, lon, result in points: - parent.add_child(make_point(lat, lon, result)) + parent.add_child(point_factory(lat, lon, result)) return branch @@ -121,3 +177,92 @@ def reforecast_tree(branches, date=np.datetime64("2024-03-01")): for b in branches: root.add_child(b) return tree + + +try: + # Only available on polytope versions that return array-backed bulk lat/lon leaves. + from polytope_feature.datacube.tensor_index_tree import ( + BulkGridTensorIndexNode, + BulkMergedTensorIndexNode, + ) +except ImportError: + BulkGridTensorIndexNode = None + BulkMergedTensorIndexNode = None + + +def _latlon_axes(): + lat_ax = IntDatacubeAxis() + lat_ax.name = "latitude" + lon_ax = IntDatacubeAxis() + lon_ax.name = "longitude" + return lat_ax, lon_ax + + +def _set_bulk_result(bulk, point_results): + """Store per-point result lists as one array of all points per combination, as polytope does.""" + n_combos = len(point_results[0]) + bulk.result = [np.array([res[c] for res in point_results], dtype=np.float64) for c in range(n_combos)] + + +def bulkify(tree): + """Replace, under every node, the compacted ``MergedTensorIndexNode`` children by one bulk leaf. + + Turns a tree built with :func:`make_merged_point` into the layout polytope returns + for unstructured grids. + """ + children = list(tree.children) + if children and all(isinstance(c, MergedTensorIndexNode) for c in children): + coordinates = [c.values for c in children] + bulk = BulkMergedTensorIndexNode(_latlon_axes(), coordinates, list(range(len(children)))) + _set_bulk_result(bulk, [list(c.result) for c in children]) + for c in children: + tree.children.remove(c) + tree.add_child(bulk) + return tree + for c in children: + bulkify(c) + return tree + + +def gridify(tree): + """Replace, under every node, the ``latitude -> longitude`` leaf layers by one bulk grid leaf. + + Turns a legacy tree (eg. built with :func:`make_point` or :func:`make_row`) into the + layout polytope returns for structured grids with ``bulk_grid_leaves``. + """ + children = list(tree.children) + if children and all(getattr(c, "axis", None) is not None and c.axis.name == "latitude" for c in children): + lat_values, lon_rows, point_results = [], [], [] + for lat_node in children: + row = [] + for lon_leaf in lat_node.children: + n = len(lon_leaf.values) + n_combos = len(lon_leaf.result) // n + for j, lon in enumerate(lon_leaf.values): + row.append(lon) + point_results.append([lon_leaf.result[c * n + j] for c in range(n_combos)]) + lat_values.append(lat_node.values[0]) + lon_rows.append(row) + grid = BulkGridTensorIndexNode(_latlon_axes(), lat_values, lon_rows) + _set_bulk_result(grid, point_results) + for c in children: + tree.children.remove(c) + tree.add_child(grid) + return tree + for c in children: + gridify(c) + return tree + + +def make_row(lat, lons, point_results): + """Create a legacy latitude->longitude(leaf) subtree holding several compressed longitudes. + + ``point_results[j]`` holds the per-combination values of longitude ``j``; the leaf's + result is laid out combination-major, as polytope assigns it. + """ + lat_n = node("latitude", (lat,)) + leaf = node("longitude", tuple(lons)) + n_combos = len(point_results[0]) + leaf.result = [np.float64(point_results[j][c]) for c in range(n_combos) for j in range(len(lons))] + lat_n.add_child(leaf) + return lat_n diff --git a/tests/test_encoder_bulk_node_parity.py b/tests/test_encoder_bulk_node_parity.py new file mode 100644 index 0000000..5a69de1 --- /dev/null +++ b/tests/test_encoder_bulk_node_parity.py @@ -0,0 +1,166 @@ +"""Parity tests for array-backed bulk lat/lon results. + +Polytope returns all the (latitude, longitude) points under one path as a single +array-backed leaf: a ``BulkMergedTensorIndexNode`` for unstructured grids, or a +``BulkGridTensorIndexNode`` (which keeps the latitude rows and their compressed +longitudes) for structured grids. These tests assert that covjsonkit produces +byte-identical CoverageJSON whether the source tree uses the legacy layout or a +bulk leaf, exercising all tree walkers. +""" + +import numpy as np +import pandas as pd +import pytest +from conftest import ( + BulkMergedTensorIndexNode, + bulkify, + chain, + forecast_tree, + gridify, + make_merged_point, + make_point, + make_row, + month_tree, + node, + reforecast_branch, + reforecast_tree, + tip, +) +from polytope_feature.datacube.tensor_index_tree import TensorIndexTree + +from covjsonkit.api import Covjsonkit + +pytestmark = pytest.mark.skipif( + BulkMergedTensorIndexNode is None, + reason="polytope build lacks BulkMergedTensorIndexNode (array-backed lat/lon leaves)", +) + +TWO_POINTS = [(48.0, 11.0, [264.9]), (50.0, 12.0, [265.1])] +MULTI_STEP_POINTS = [(48.0, 11.0, [264.9, 270.1]), (50.0, 12.0, [265.1, 271.3]), (50.0, 13.0, [266.0, 272.0])] + +FEATURES = ["BoundingBox", "Grid", "Frame", "Circle", "Shapefile", "Polygon", "Path", "Position"] + + +def _encode(tree, feature="BoundingBox", reforecast=False): + api = Covjsonkit().encode("CoverageCollection", feature) + if reforecast: + return api.from_polytope_reforecast(tree) + return api.from_polytope(tree) + + +def _encode_month(tree, feature="BoundingBox"): + return Covjsonkit().encode("CoverageCollection", feature).from_polytope_month(tree) + + +def _encode_step(tree, feature="PointSeries"): + return Covjsonkit().encode("CoverageCollection", feature).from_polytope_step(tree) + + +class TestBulkNodeParity: + @pytest.mark.parametrize("feature", FEATURES) + def test_walk_tree_single_step(self, feature): + legacy = _encode(forecast_tree(TWO_POINTS, point_factory=make_point), feature) + bulk = _encode(bulkify(forecast_tree(TWO_POINTS, point_factory=make_merged_point)), feature) + grid = _encode(gridify(forecast_tree(TWO_POINTS, point_factory=make_point)), feature) + assert bulk == legacy + assert grid == legacy + + def test_walk_tree_multi_step(self): + legacy = _encode(forecast_tree(MULTI_STEP_POINTS, step=(0, 6), point_factory=make_point)) + bulk = _encode(bulkify(forecast_tree(MULTI_STEP_POINTS, step=(0, 6), point_factory=make_merged_point))) + grid = _encode(gridify(forecast_tree(MULTI_STEP_POINTS, step=(0, 6), point_factory=make_point))) + assert bulk == legacy + assert grid == legacy + + def test_grid_row_with_compressed_longitudes(self): + # One latitude row holding two compressed longitudes, as the hullslicer returns them + def build(): + tree = forecast_tree([], step=(0, 6)) + fc = tip(tree) + fc.add_child(make_row(48.0, (11.0, 11.5), [[264.9, 270.1], [264.5, 270.5]])) + fc.add_child(make_row(50.0, (12.0,), [[265.1, 271.3]])) + return tree + + legacy = _encode(build()) + grid = _encode(gridify(build())) + assert len(legacy["coverages"]) == 2 + assert grid == legacy + + def test_walk_tree_two_dates_two_steps(self): + def build(point_factory): + tree = chain(TensorIndexTree(), node("class", ("od",))) + cls = tip(tree) + for date_val, vals in [ + (np.datetime64("2025-01-01T00:00:00"), [[264.9, 270.1], [265.1, 271.3]]), + (np.datetime64("2025-01-02T00:00:00"), [[266.0, 272.0], [267.0, 273.0]]), + ]: + branch = chain( + node("date", (date_val,)), + node("domain", ("g",)), + node("expver", ("0001",)), + node("levtype", ("sfc",)), + node("param", ("167",)), + node("step", (0, 6)), + node("stream", ("oper",)), + node("type", ("fc",)), + ) + fc = tip(branch) + fc.add_child(point_factory(48.0, 11.0, vals[0])) + fc.add_child(point_factory(50.0, 12.0, vals[1])) + cls.add_child(branch) + return tree + + legacy = _encode(build(make_point)) + assert _encode(bulkify(build(make_merged_point))) == legacy + assert _encode(gridify(build(make_point))) == legacy + + def test_reforecast_walker(self): + def build(point_factory): + return reforecast_tree( + [ + reforecast_branch(np.datetime64("2025-07-14T06:00:00"), TWO_POINTS, point_factory=point_factory), + reforecast_branch( + np.datetime64("2025-07-15T06:00:00"), + [(48.0, 11.0, [266.0]), (50.0, 12.0, [267.0])], + point_factory=point_factory, + ), + ] + ) + + legacy = _encode(build(make_point), reforecast=True) + assert _encode(bulkify(build(make_merged_point)), reforecast=True) == legacy + assert _encode(gridify(build(make_point)), reforecast=True) == legacy + + def test_walk_tree_month(self): + points = [(48.0, 11.0, [264.9, 265.9]), (50.0, 12.0, [266.1, 267.1])] + legacy = _encode_month(month_tree(points, point_factory=make_point)) + bulk = _encode_month(bulkify(month_tree(points, point_factory=make_merged_point))) + grid = _encode_month(gridify(month_tree(points, point_factory=make_point))) + assert len(legacy["coverages"]) == 2 + assert bulk == legacy + assert grid == legacy + + def test_walk_tree_step(self): + # Time series (step walker): date -> time axes, two points over two times + def build(point_factory): + tree = chain( + TensorIndexTree(), + node("class", ("od",)), + node("date", (np.datetime64("2025-01-01T00:00:00"),)), + node("time", (pd.Timedelta(hours=0), pd.Timedelta(hours=6))), + node("domain", ("g",)), + node("expver", ("0001",)), + node("levtype", ("sfc",)), + node("param", ("167",)), + node("stream", ("oper",)), + node("type", ("fc",)), + ) + fc = tip(tree) + fc.add_child(point_factory(48.0, 11.0, [264.9, 270.1])) + fc.add_child(point_factory(50.0, 12.0, [265.1, 271.3])) + return tree + + legacy = _encode_step(build(make_point)) + assert len(legacy["coverages"]) == 2 + assert _encode_step(bulkify(build(make_merged_point))) == legacy + assert _encode_step(gridify(build(make_point))) == legacy diff --git a/tests/test_encoder_merged_node_parity.py b/tests/test_encoder_merged_node_parity.py new file mode 100644 index 0000000..feb21a7 --- /dev/null +++ b/tests/test_encoder_merged_node_parity.py @@ -0,0 +1,111 @@ +"""Parity tests for compacted unstructured-grid results. + +Polytope emits ``MergedTensorIndexNode`` leaves for unstructured grids (ICON, +Lambert LAM), compacting each (latitude, longitude) point and its result into a +single leaf instead of a latitude -> longitude-leaf subtree. These tests assert +that covjsonkit produces byte-identical CoverageJSON whether the source tree +uses the legacy layout (``make_point``) or the compacted layout +(``make_merged_point``), exercising all tree walkers. +""" + +import numpy as np +import pytest +from conftest import ( + MergedTensorIndexNode, + chain, + forecast_tree, + make_merged_point, + make_point, + month_tree, + node, + reforecast_branch, + reforecast_tree, + tip, +) +from polytope_feature.datacube.tensor_index_tree import TensorIndexTree + +from covjsonkit.api import Covjsonkit + +# These tests require a polytope build with compacted unstructured-grid support. +# On released polytope (no MergedTensorIndexNode) they are skipped entirely. +pytestmark = pytest.mark.skipif( + MergedTensorIndexNode is None, + reason="polytope build lacks MergedTensorIndexNode (compacted unstructured support)", +) + +TWO_POINTS = [(48.0, 11.0, [264.9]), (50.0, 12.0, [265.1])] + + +def _encode(tree, feature="BoundingBox", reforecast=False): + api = Covjsonkit().encode("CoverageCollection", feature) + if reforecast: + return api.from_polytope_reforecast(tree) + return api.from_polytope(tree) + + +def _encode_month(tree, feature="BoundingBox"): + return Covjsonkit().encode("CoverageCollection", feature).from_polytope_month(tree) + + +class TestMergedNodeParity: + def test_walk_tree_single_date_single_step(self): + legacy = _encode(forecast_tree(TWO_POINTS, point_factory=make_point)) + merged = _encode(forecast_tree(TWO_POINTS, point_factory=make_merged_point)) + assert merged == legacy + + def test_walk_tree_multi_step(self): + points = [(48.0, 11.0, [264.9, 270.1]), (50.0, 12.0, [265.1, 271.3])] + legacy = _encode(forecast_tree(points, step=(0, 6), point_factory=make_point)) + merged = _encode(forecast_tree(points, step=(0, 6), point_factory=make_merged_point)) + assert merged == legacy + + def test_walk_tree_two_dates_two_steps(self): + def build(point_factory): + tree = chain(TensorIndexTree(), node("class", ("od",))) + cls = tip(tree) + for date_val, vals in [ + (np.datetime64("2025-01-01T00:00:00"), [[264.9, 270.1], [265.1, 271.3]]), + (np.datetime64("2025-01-02T00:00:00"), [[266.0, 272.0], [267.0, 273.0]]), + ]: + branch = chain( + node("date", (date_val,)), + node("domain", ("g",)), + node("expver", ("0001",)), + node("levtype", ("sfc",)), + node("param", ("167",)), + node("step", (0, 6)), + node("stream", ("oper",)), + node("type", ("fc",)), + ) + fc = tip(branch) + fc.add_child(point_factory(48.0, 11.0, vals[0])) + fc.add_child(point_factory(50.0, 12.0, vals[1])) + cls.add_child(branch) + return tree + + assert _encode(build(make_merged_point)) == _encode(build(make_point)) + + def test_reforecast_walker(self): + def build(point_factory): + return reforecast_tree( + [ + reforecast_branch(np.datetime64("2025-07-14T06:00:00"), TWO_POINTS, point_factory=point_factory), + reforecast_branch( + np.datetime64("2025-07-15T06:00:00"), + [(48.0, 11.0, [266.0]), (50.0, 12.0, [267.0])], + point_factory=point_factory, + ), + ] + ) + + legacy = _encode(build(make_point), reforecast=True) + merged = _encode(build(make_merged_point), reforecast=True) + assert merged == legacy + + def test_walk_tree_month(self): + # Monthly-mean (year/month axes) path: two years x one month, two points. + points = [(48.0, 11.0, [264.9, 265.9]), (50.0, 12.0, [266.1, 267.1])] + legacy = _encode_month(month_tree(points, point_factory=make_point)) + merged = _encode_month(month_tree(points, point_factory=make_merged_point)) + assert len(legacy["coverages"]) == 2 + assert merged == legacy