Loading csst_dag/v2/dag/__init__.py +11 −3 Original line number Diff line number Diff line from .base import BaseDAG, DependencyBundle from .brick import BrickDAG from .brick import BrickDAG, CatalogBrickDAG, FileBrickDAG from .dispatcher import Dispatcher from .file import FileDAG Loading Loading @@ -54,10 +54,16 @@ CSST_DAGS = { "csst-ifs-l2", dispatcher=Dispatcher.dispatch_obsid, ), "csst-msc-l2-mbi-mosaic": BrickDAG( "csst-msc-l2-mbi-mosaic": FileBrickDAG( "csst-msc-l2-mbi-mosaic", ), "csst-msc-l2-mbi-xcat": BrickDAG( "csst-msc-l2-mbi-fphot": FileBrickDAG( "csst-msc-l2-mbi-fphot", ), "csst-msc-l2-mbi-photoz": FileBrickDAG( "csst-msc-l2-mbi-photoz", ), "csst-msc-l2-mbi-xcat": CatalogBrickDAG( "csst-msc-l2-mbi-xcat", ), } Loading @@ -65,8 +71,10 @@ CSST_DAGS = { __all__ = [ "BaseDAG", "BrickDAG", "CatalogBrickDAG", "CSST_DAGS", "DependencyBundle", "Dispatcher", "FileDAG", "FileBrickDAG", ] csst_dag/v2/dag/brick.py +159 −74 Original line number Diff line number Diff line from __future__ import annotations from itertools import product from typing import Any, Optional from astropy.table import Table from astropy.table import Column, Table from ..config_loader import load_dag_config from ..dfs import data as dfs_data from ..dfs.model_registry import get_storage_level from ..models import DagRun, DagRunGroup from .base import BaseDAG, DependencyBundle class BrickDAG(BaseDAG): class BaseBrickDAG(BaseDAG): """ v2 Brick DAG。 v2 Brick DAG 基类。 Notes ----- Brick DAG 的触发以 `healpix` 为分组键。 当前逻辑直接复用 submission 页面里的 level1 查询条件, 先筛出 level1 数据,再按唯一 `healpix` 生成一批 DagRun。 `data_model` 视为本次 submission 的数据检索条件之一, 需要由调用方在 `data` 中显式传入,且必须为 level1 data model。 Brick DAG 通过 submission 页面 `data` 面板中的筛选条件, 调用 DFS 的 `find_unique_healpix()` 检索唯一 healpix 列表。 随后以 healpix 为基本分组键生成 DagRun;若 DAG 配置中声明了 `brick_multipliers`,则进一步生成 `healpix × multiplier` 的排列组合。 """ source_level = "" UNIQUE_QUERY_KEYS = ( "data_model", "dataset", "batch_id", "obs_group", "obs_id", "instrument", "detector", "obs_type", "filter", "healpix", "custom_id", "prc_status", "qc_status", ) def __init__(self, dag_name: str): super().__init__(dag_name=dag_name, require_plan=False, require_data=True, require_upstream=True) self.dag_cfg = load_dag_config(dag_name) self.brick_multipliers_cfg = self.dag_cfg.get("brick_multipliers") or [] def build_dependency_queries( self, Loading @@ -33,31 +54,15 @@ class BrickDAG(BaseDAG): plan_query, data_query, upstream_query = super().build_dependency_queries(data, proc) data_model = str(data_query.get("data_model") or "").strip() if not data_model or data_model.lower() == "raw": raise ValueError("BrickDAG 需要在 data 中显式提供非 raw 的 data_model。") return plan_query, data_query, upstream_query def resolve_data_dependency(self, data_query: dict[str, Any]) -> Table: return dfs_data.find(**data_query) raise ValueError(f"{self.__class__.__name__} 需要在 data 中显式提供非 raw 的 data_model。") @staticmethod def _pick_group_value(rows: Table, column: str, preferred: Any = None) -> Optional[str]: """优先使用用户显式输入,否则仅在组内唯一时回填该字段。""" if preferred not in (None, ""): return str(preferred) if column not in rows.colnames: return None values: list[str] = [] seen: set[str] = set() for value in rows[column]: if value in (None, ""): continue text = str(value) if text in seen: continue seen.add(text) values.append(text) return values[0] if len(values) == 1 else None actual_level = get_storage_level(data_model) if actual_level != self.source_level: raise ValueError( f"{self.__class__.__name__} 需要 `{self.source_level}` data_model," f"当前 `{data_model}` 注册在 `{actual_level}`。" ) return plan_query, data_query, upstream_query @staticmethod def _normalize_healpix(value: Any) -> Optional[str]: Loading @@ -69,23 +74,79 @@ class BrickDAG(BaseDAG): return None return text def _build_unique_healpix_query(self, data_query: dict[str, Any]) -> dict[str, Any]: """仅保留 unique_healpix 接口需要的筛选条件。""" return { key: value for key, value in data_query.items() if key in self.UNIQUE_QUERY_KEYS and value not in (None, "") } def resolve_unique_healpix(self, data_query: dict[str, Any]) -> list[str]: """统一通过 v2 DFS wrapper 查询唯一 healpix。""" return dfs_data.find_unique_healpix(**data_query) def resolve_data_dependency(self, data_query: dict[str, Any]) -> Table: unique_healpix = self.resolve_unique_healpix(self._build_unique_healpix_query(data_query)) if not unique_healpix: return Table([Column([], dtype=object, name="healpix")]) return Table([{"healpix": value} for value in unique_healpix]) @staticmethod def _iter_unique_healpix(rows: Table) -> list[str]: """保持原始顺序提取唯一的整数 healpix。""" if "healpix" not in rows.colnames: def _normalize_multiplier_values(values: Any) -> list[str]: """规范化 multiplier 的 values 列表并去重。""" if values in (None, ""): return [] if not isinstance(values, (list, tuple, set)): values = [values] values: list[str] = [] normalized: list[str] = [] seen: set[str] = set() for value in rows["healpix"]: text = BrickDAG._normalize_healpix(value) if text is None: for item in values: if item in (None, ""): continue text = str(item) if text in seen: continue seen.add(text) values.append(text) return values normalized.append(text) return normalized def _resolve_run_multiplier_values(self, data: dict[str, Any], field: str, values: Any) -> list[str]: """优先使用 submission 显式值,否则回退到 DAG 配置中的 multiplier。""" explicit_value = data.get(field) if explicit_value not in (None, ""): return [str(explicit_value)] return self._normalize_multiplier_values(values) def _resolve_run_multiplier_combinations(self, data: dict[str, Any]) -> list[dict[str, str]]: """解析并展开本次 Brick 需要附加到 DagRun 的 multiplier 组合。""" if not isinstance(self.brick_multipliers_cfg, list): raise ValueError("`brick_multipliers` 必须是列表。") specs: list[tuple[str, list[str]]] = [] for item in self.brick_multipliers_cfg: if not isinstance(item, dict): raise ValueError("`brick_multipliers` 的每一项都必须是字典。") field = str(item.get("field") or "").strip() if not field: raise ValueError("`brick_multipliers` 缺少 `field`。") values = self._resolve_run_multiplier_values(data, field, item.get("values")) if not values: return [] specs.append((field, values)) if not specs: return [{}] combinations: list[dict[str, str]] = [] for combo_values in product(*[values for _, values in specs]): combinations.append({ field: value for (field, _), value in zip(specs, combo_values) }) return combinations def build_trigger_payload( self, Loading @@ -104,39 +165,63 @@ class BrickDAG(BaseDAG): if deps.data is None or len(deps.data) == 0: return dag_group unique_healpix = self._iter_unique_healpix(deps.data) for this_healpix in unique_healpix: rows = deps.data[ [self._normalize_healpix(v) == this_healpix for v in deps.data["healpix"]] run_multiplier_combinations = self._resolve_run_multiplier_combinations(data) if not run_multiplier_combinations: return dag_group unique_healpix = [ self._normalize_healpix(value) for value in deps.data["healpix"] ] run = DagRun( dag=self.dag_name, dag_run_group=dag_group["dag_run_group"], batch_id=dag_group.get("batch_id"), priority=dag_group.get("priority", "low"), data_model=self._pick_group_value(rows, "data_model", data.get("data_model")), dataset=self._pick_group_value(rows, "dataset", data.get("dataset")), instrument=self._pick_group_value(rows, "instrument", data.get("instrument")), obs_type=self._pick_group_value(rows, "obs_type", data.get("obs_type")), obs_group=self._pick_group_value(rows, "obs_group", data.get("obs_group")), obs_id=self._pick_group_value(rows, "obs_id", data.get("obs_id")), detector=self._pick_group_value(rows, "detector", data.get("detector")), filter=self._pick_group_value(rows, "filter", data.get("filter")), healpix=this_healpix, custom_id=data.get("custom_id"), prc_status=data.get("prc_status"), qc_status=data.get("qc_status"), n_file_expected=len(rows), n_file_found=len(rows), data_list=( [str(v) for v in rows["data_uuid"] if v not in (None, "")] if "data_uuid" in rows.colnames else [] base_run_kwargs = { "dag": self.dag_name, "dag_run_group": dag_group["dag_run_group"], "batch_id": dag_group.get("batch_id"), "priority": dag_group.get("priority", "low"), "data_model": (str(data.get("data_model")) if data.get("data_model") not in (None, "") else None), "dataset": (str(data.get("dataset")) if data.get("dataset") not in (None, "") else None), "source_batch_id": ( str(data.get("source_batch_id")) if data.get("source_batch_id") not in (None, "") else None ), pmapname=proc.get("pmapname"), ref_cat=proc.get("ref_cat"), extra_kwargs=proc.get("extra_kwargs"), ) "instrument": (str(data.get("instrument")) if data.get("instrument") not in (None, "") else None), "obs_type": (str(data.get("obs_type")) if data.get("obs_type") not in (None, "") else None), "obs_group": (str(data.get("obs_group")) if data.get("obs_group") not in (None, "") else None), "obs_id": (str(data.get("obs_id")) if data.get("obs_id") not in (None, "") else None), "detector": (str(data.get("detector")) if data.get("detector") not in (None, "") else None), "filter": (str(data.get("filter")) if data.get("filter") not in (None, "") else None), "custom_id": (str(data.get("custom_id")) if data.get("custom_id") not in (None, "") else None), "prc_status": data.get("prc_status"), "qc_status": data.get("qc_status"), "pmapname": proc.get("pmapname"), "ref_cat": proc.get("ref_cat"), "extra_kwargs": proc.get("extra_kwargs"), } for this_healpix in unique_healpix: if this_healpix is None: continue for multiplier_values in run_multiplier_combinations: run_kwargs = dict(base_run_kwargs) run_kwargs.update(multiplier_values) run_kwargs["healpix"] = this_healpix run = DagRun(**run_kwargs) dag_group.append_dag_run(run) return dag_group class FileBrickDAG(BaseBrickDAG): """基于 `level1.find_unique_healpix()` 的文件型 Brick DAG。""" source_level = "level1" class CatalogBrickDAG(BaseBrickDAG): """基于 `level2.find_unique_healpix()` 的星表型 Brick DAG。""" source_level = "level2" class BrickDAG(FileBrickDAG): """兼容旧导入路径的别名,默认仍表示 level1 文件型 Brick。""" csst_dag/v2/dag/file.py +1 −0 Original line number Diff line number Diff line Loading @@ -127,6 +127,7 @@ class FileDAG(BaseDAG): priority=dag_group.get("priority", "low"), data_model=task.get("data_model"), dataset=task.get("dataset"), source_batch_id=data.get("source_batch_id"), instrument=task.get("instrument"), obs_type=task.get("obs_type"), obs_group=task.get("obs_group"), Loading csst_dag/v2/dag_config/csst-cpic-l1-qc0.yml +1 −0 Original line number Diff line number Diff line Loading @@ -46,4 +46,5 @@ tasks: - name: QC0 image: csst-cpic-l1-qc0 command: ["run", "{{ params.input }}"] pool_slots: 1 csst_dag/v2/dag_config/csst-cpic-l1.yml +1 −0 Original line number Diff line number Diff line Loading @@ -47,4 +47,5 @@ tasks: - name: CPIC image: csst-cpic-l1 command: ["run", "{{ params.input }}"] pool_slots: 1 Loading
csst_dag/v2/dag/__init__.py +11 −3 Original line number Diff line number Diff line from .base import BaseDAG, DependencyBundle from .brick import BrickDAG from .brick import BrickDAG, CatalogBrickDAG, FileBrickDAG from .dispatcher import Dispatcher from .file import FileDAG Loading Loading @@ -54,10 +54,16 @@ CSST_DAGS = { "csst-ifs-l2", dispatcher=Dispatcher.dispatch_obsid, ), "csst-msc-l2-mbi-mosaic": BrickDAG( "csst-msc-l2-mbi-mosaic": FileBrickDAG( "csst-msc-l2-mbi-mosaic", ), "csst-msc-l2-mbi-xcat": BrickDAG( "csst-msc-l2-mbi-fphot": FileBrickDAG( "csst-msc-l2-mbi-fphot", ), "csst-msc-l2-mbi-photoz": FileBrickDAG( "csst-msc-l2-mbi-photoz", ), "csst-msc-l2-mbi-xcat": CatalogBrickDAG( "csst-msc-l2-mbi-xcat", ), } Loading @@ -65,8 +71,10 @@ CSST_DAGS = { __all__ = [ "BaseDAG", "BrickDAG", "CatalogBrickDAG", "CSST_DAGS", "DependencyBundle", "Dispatcher", "FileDAG", "FileBrickDAG", ]
csst_dag/v2/dag/brick.py +159 −74 Original line number Diff line number Diff line from __future__ import annotations from itertools import product from typing import Any, Optional from astropy.table import Table from astropy.table import Column, Table from ..config_loader import load_dag_config from ..dfs import data as dfs_data from ..dfs.model_registry import get_storage_level from ..models import DagRun, DagRunGroup from .base import BaseDAG, DependencyBundle class BrickDAG(BaseDAG): class BaseBrickDAG(BaseDAG): """ v2 Brick DAG。 v2 Brick DAG 基类。 Notes ----- Brick DAG 的触发以 `healpix` 为分组键。 当前逻辑直接复用 submission 页面里的 level1 查询条件, 先筛出 level1 数据,再按唯一 `healpix` 生成一批 DagRun。 `data_model` 视为本次 submission 的数据检索条件之一, 需要由调用方在 `data` 中显式传入,且必须为 level1 data model。 Brick DAG 通过 submission 页面 `data` 面板中的筛选条件, 调用 DFS 的 `find_unique_healpix()` 检索唯一 healpix 列表。 随后以 healpix 为基本分组键生成 DagRun;若 DAG 配置中声明了 `brick_multipliers`,则进一步生成 `healpix × multiplier` 的排列组合。 """ source_level = "" UNIQUE_QUERY_KEYS = ( "data_model", "dataset", "batch_id", "obs_group", "obs_id", "instrument", "detector", "obs_type", "filter", "healpix", "custom_id", "prc_status", "qc_status", ) def __init__(self, dag_name: str): super().__init__(dag_name=dag_name, require_plan=False, require_data=True, require_upstream=True) self.dag_cfg = load_dag_config(dag_name) self.brick_multipliers_cfg = self.dag_cfg.get("brick_multipliers") or [] def build_dependency_queries( self, Loading @@ -33,31 +54,15 @@ class BrickDAG(BaseDAG): plan_query, data_query, upstream_query = super().build_dependency_queries(data, proc) data_model = str(data_query.get("data_model") or "").strip() if not data_model or data_model.lower() == "raw": raise ValueError("BrickDAG 需要在 data 中显式提供非 raw 的 data_model。") return plan_query, data_query, upstream_query def resolve_data_dependency(self, data_query: dict[str, Any]) -> Table: return dfs_data.find(**data_query) raise ValueError(f"{self.__class__.__name__} 需要在 data 中显式提供非 raw 的 data_model。") @staticmethod def _pick_group_value(rows: Table, column: str, preferred: Any = None) -> Optional[str]: """优先使用用户显式输入,否则仅在组内唯一时回填该字段。""" if preferred not in (None, ""): return str(preferred) if column not in rows.colnames: return None values: list[str] = [] seen: set[str] = set() for value in rows[column]: if value in (None, ""): continue text = str(value) if text in seen: continue seen.add(text) values.append(text) return values[0] if len(values) == 1 else None actual_level = get_storage_level(data_model) if actual_level != self.source_level: raise ValueError( f"{self.__class__.__name__} 需要 `{self.source_level}` data_model," f"当前 `{data_model}` 注册在 `{actual_level}`。" ) return plan_query, data_query, upstream_query @staticmethod def _normalize_healpix(value: Any) -> Optional[str]: Loading @@ -69,23 +74,79 @@ class BrickDAG(BaseDAG): return None return text def _build_unique_healpix_query(self, data_query: dict[str, Any]) -> dict[str, Any]: """仅保留 unique_healpix 接口需要的筛选条件。""" return { key: value for key, value in data_query.items() if key in self.UNIQUE_QUERY_KEYS and value not in (None, "") } def resolve_unique_healpix(self, data_query: dict[str, Any]) -> list[str]: """统一通过 v2 DFS wrapper 查询唯一 healpix。""" return dfs_data.find_unique_healpix(**data_query) def resolve_data_dependency(self, data_query: dict[str, Any]) -> Table: unique_healpix = self.resolve_unique_healpix(self._build_unique_healpix_query(data_query)) if not unique_healpix: return Table([Column([], dtype=object, name="healpix")]) return Table([{"healpix": value} for value in unique_healpix]) @staticmethod def _iter_unique_healpix(rows: Table) -> list[str]: """保持原始顺序提取唯一的整数 healpix。""" if "healpix" not in rows.colnames: def _normalize_multiplier_values(values: Any) -> list[str]: """规范化 multiplier 的 values 列表并去重。""" if values in (None, ""): return [] if not isinstance(values, (list, tuple, set)): values = [values] values: list[str] = [] normalized: list[str] = [] seen: set[str] = set() for value in rows["healpix"]: text = BrickDAG._normalize_healpix(value) if text is None: for item in values: if item in (None, ""): continue text = str(item) if text in seen: continue seen.add(text) values.append(text) return values normalized.append(text) return normalized def _resolve_run_multiplier_values(self, data: dict[str, Any], field: str, values: Any) -> list[str]: """优先使用 submission 显式值,否则回退到 DAG 配置中的 multiplier。""" explicit_value = data.get(field) if explicit_value not in (None, ""): return [str(explicit_value)] return self._normalize_multiplier_values(values) def _resolve_run_multiplier_combinations(self, data: dict[str, Any]) -> list[dict[str, str]]: """解析并展开本次 Brick 需要附加到 DagRun 的 multiplier 组合。""" if not isinstance(self.brick_multipliers_cfg, list): raise ValueError("`brick_multipliers` 必须是列表。") specs: list[tuple[str, list[str]]] = [] for item in self.brick_multipliers_cfg: if not isinstance(item, dict): raise ValueError("`brick_multipliers` 的每一项都必须是字典。") field = str(item.get("field") or "").strip() if not field: raise ValueError("`brick_multipliers` 缺少 `field`。") values = self._resolve_run_multiplier_values(data, field, item.get("values")) if not values: return [] specs.append((field, values)) if not specs: return [{}] combinations: list[dict[str, str]] = [] for combo_values in product(*[values for _, values in specs]): combinations.append({ field: value for (field, _), value in zip(specs, combo_values) }) return combinations def build_trigger_payload( self, Loading @@ -104,39 +165,63 @@ class BrickDAG(BaseDAG): if deps.data is None or len(deps.data) == 0: return dag_group unique_healpix = self._iter_unique_healpix(deps.data) for this_healpix in unique_healpix: rows = deps.data[ [self._normalize_healpix(v) == this_healpix for v in deps.data["healpix"]] run_multiplier_combinations = self._resolve_run_multiplier_combinations(data) if not run_multiplier_combinations: return dag_group unique_healpix = [ self._normalize_healpix(value) for value in deps.data["healpix"] ] run = DagRun( dag=self.dag_name, dag_run_group=dag_group["dag_run_group"], batch_id=dag_group.get("batch_id"), priority=dag_group.get("priority", "low"), data_model=self._pick_group_value(rows, "data_model", data.get("data_model")), dataset=self._pick_group_value(rows, "dataset", data.get("dataset")), instrument=self._pick_group_value(rows, "instrument", data.get("instrument")), obs_type=self._pick_group_value(rows, "obs_type", data.get("obs_type")), obs_group=self._pick_group_value(rows, "obs_group", data.get("obs_group")), obs_id=self._pick_group_value(rows, "obs_id", data.get("obs_id")), detector=self._pick_group_value(rows, "detector", data.get("detector")), filter=self._pick_group_value(rows, "filter", data.get("filter")), healpix=this_healpix, custom_id=data.get("custom_id"), prc_status=data.get("prc_status"), qc_status=data.get("qc_status"), n_file_expected=len(rows), n_file_found=len(rows), data_list=( [str(v) for v in rows["data_uuid"] if v not in (None, "")] if "data_uuid" in rows.colnames else [] base_run_kwargs = { "dag": self.dag_name, "dag_run_group": dag_group["dag_run_group"], "batch_id": dag_group.get("batch_id"), "priority": dag_group.get("priority", "low"), "data_model": (str(data.get("data_model")) if data.get("data_model") not in (None, "") else None), "dataset": (str(data.get("dataset")) if data.get("dataset") not in (None, "") else None), "source_batch_id": ( str(data.get("source_batch_id")) if data.get("source_batch_id") not in (None, "") else None ), pmapname=proc.get("pmapname"), ref_cat=proc.get("ref_cat"), extra_kwargs=proc.get("extra_kwargs"), ) "instrument": (str(data.get("instrument")) if data.get("instrument") not in (None, "") else None), "obs_type": (str(data.get("obs_type")) if data.get("obs_type") not in (None, "") else None), "obs_group": (str(data.get("obs_group")) if data.get("obs_group") not in (None, "") else None), "obs_id": (str(data.get("obs_id")) if data.get("obs_id") not in (None, "") else None), "detector": (str(data.get("detector")) if data.get("detector") not in (None, "") else None), "filter": (str(data.get("filter")) if data.get("filter") not in (None, "") else None), "custom_id": (str(data.get("custom_id")) if data.get("custom_id") not in (None, "") else None), "prc_status": data.get("prc_status"), "qc_status": data.get("qc_status"), "pmapname": proc.get("pmapname"), "ref_cat": proc.get("ref_cat"), "extra_kwargs": proc.get("extra_kwargs"), } for this_healpix in unique_healpix: if this_healpix is None: continue for multiplier_values in run_multiplier_combinations: run_kwargs = dict(base_run_kwargs) run_kwargs.update(multiplier_values) run_kwargs["healpix"] = this_healpix run = DagRun(**run_kwargs) dag_group.append_dag_run(run) return dag_group class FileBrickDAG(BaseBrickDAG): """基于 `level1.find_unique_healpix()` 的文件型 Brick DAG。""" source_level = "level1" class CatalogBrickDAG(BaseBrickDAG): """基于 `level2.find_unique_healpix()` 的星表型 Brick DAG。""" source_level = "level2" class BrickDAG(FileBrickDAG): """兼容旧导入路径的别名,默认仍表示 level1 文件型 Brick。"""
csst_dag/v2/dag/file.py +1 −0 Original line number Diff line number Diff line Loading @@ -127,6 +127,7 @@ class FileDAG(BaseDAG): priority=dag_group.get("priority", "low"), data_model=task.get("data_model"), dataset=task.get("dataset"), source_batch_id=data.get("source_batch_id"), instrument=task.get("instrument"), obs_type=task.get("obs_type"), obs_group=task.get("obs_group"), Loading
csst_dag/v2/dag_config/csst-cpic-l1-qc0.yml +1 −0 Original line number Diff line number Diff line Loading @@ -46,4 +46,5 @@ tasks: - name: QC0 image: csst-cpic-l1-qc0 command: ["run", "{{ params.input }}"] pool_slots: 1
csst_dag/v2/dag_config/csst-cpic-l1.yml +1 −0 Original line number Diff line number Diff line Loading @@ -47,4 +47,5 @@ tasks: - name: CPIC image: csst-cpic-l1 command: ["run", "{{ params.input }}"] pool_slots: 1