Commit edd8a5a2 authored by BO ZHANG's avatar BO ZHANG 🏀
Browse files

add CALF

parent a2d3d76e
Loading
Loading
Loading
Loading
+18 −2
Original line number Diff line number Diff line
from ._base_dag import BaseDAG, Level1DAG, Level2DAG
from ._base_dag import BaseDAG, Level1DAG, Level2DAG, BrickDAG
from ._dispatcher import Dispatcher
from .._csst import csst
from ..dag_utils import generate_permutations
@@ -17,7 +17,7 @@ CSST_DAGS = {
        dag="csst-msc-l1-mbi",
        pattern=generate_permutations(
            instrument=["MSC"],
            obs_type=["WIDE", "DEEP"],
            obs_type=["WIDE", "DEEP", "CALF"],
            detector=csst["MSC"]["MBI"].effective_detector_names,
        ),
        dispatcher=Dispatcher.dispatch_file,
@@ -49,6 +49,22 @@ CSST_DAGS = {
        ),
        dispatcher=Dispatcher.dispatch_obsgroup_detector,
    ),
    # MSC brick DAGs
    "csst-msc-l2-mbi-mosaic": BrickDAG(
        dag="csst-msc-l2-mbi-mosaic",
        required_data_models=["csst-msc-l1-mbi"],
        required_dags=None,
    ),
    "csst-msc-l2-mbi-xcat": BrickDAG(
        dag="csst-msc-l2-mbi-xcat",
        required_data_models=["csst-msc-l2-mbi-cat"],
        required_dags=None,
    ),
    "csst-msc-l2-mbi-xcatmix": BrickDAG(
        dag="csst-msc-l2-mbi-xcatmix",
        required_data_models=["csst-msc-l2-mbi-xcatmix"],
        required_dags=None,
    ),
    # MCI
    "csst-mci-l1": Level1DAG(
        dag="csst-mci-l1",
+159 −0
Original line number Diff line number Diff line
@@ -369,3 +369,162 @@ class Level1DAG(BaseDAG):
            else:
                this_task["dag_run"] = None
        return task_list


class BrickDAG(BaseDAG):

    SUPPORTED_DAG_LIST = [
        "csst-msc-l2-mbi-mosaic",
        "csst-msc-l2-mbi-xcat",
        "csst-msc-l2-mbi-xcatmix",
    ]

    def __init__(
        self,
        dag: str,
        required_data_models: Optional[list[str]] = None,
        required_dags: Optional[list[str]] = None,
    ):
        super().__init__()
        assert (
            dag in self.SUPPORTED_DAG_LIST
        ), f"DAG {dag} is not supported, supported DAG list: {self.SUPPORTED_DAG_LIST}"

        self.dag = dag
        self.required_data_models = required_data_models or []
        self.required_dags = required_dags or []

    def run(
        self,
        # DAG group parameters
        dag_group: str = "default-dag-group",
        batch_id: str = "default-batch",
        priority: int | str = 1,
        # plan filter
        dataset: str | None = None,
        instrument: str | None = None,
        obs_type: str | None = None,
        obs_group: str | None = None,
        obs_id: str | None = None,
        custom_id: str | None = None,
        proposal_id: str | None = None,
        object: str | None = None,
        # data filter
        detector: str | None = None,
        filter: str | None = None,
        prc_status: str | None = None,  # no effect
        qc_status: str | None = None,
        healpix: int | None = None,
        # prc parameters
        pmapname: str = "",
        extra_kwargs: Optional[dict] = None,
        # additional parameters
        debug: bool = False,
    ) -> tuple[dict, list]:
        # if required_dags is empty, then it is a root DAG
        if not self.required_dags and self.required_data_models:

            # generate DAG group run
            dag_group_run = self.generate_dag_group_run(
                dag_group=dag_group,
                batch_id=batch_id,
                priority=priority,
            )

            # declare results: the brick ids of each data model
            data_model_brick_ids = {_: {} for _ in self.required_data_models}
            # loop over data models
            for i, this_data_model in enumerate(self.required_data_models):
                print(
                    f"Processing required data model "
                    f"[{i}/{len(self.required_data_models)}]: {this_data_model}"
                )
                # construct query filter
                query_filter = {
                    key: value
                    for key, value in {
                        "data_model": this_data_model,
                        "dataset": dataset,
                        "instrument": instrument,
                        "obs_type": obs_type,
                        "obs_group": obs_group,
                        "obs_id": obs_id,
                        "custom_id": custom_id,
                        "detector": detector,
                        "filter": filter,
                        "prc_status": prc_status,
                        "qc_status": qc_status,
                        "proposal_id": proposal_id,
                        "object": object,
                        "healpix": healpix,
                    }.items()
                    if value is not None and value != ""  # 过滤条件
                }
                query_keys = [
                    "healpix",
                ]
                # send query request
                query_results = csst_fs.query_metadata(
                    filter=query_filter,
                    key=query_keys,
                    limit=0,
                )
                if len(query_results) > 0:
                    print(f"No data found for the given constraints:{query_filter}")
                    return dag_group_run, []

                # unique healpix for this data model
                u_brick_ids = set([_["metadata"]["healpix"] for _ in query_results])

                # store brick ids for each data model
                data_model_brick_ids[this_data_model][healpix] = set(u_brick_ids)

            # calculate intersection of brick ids for each data model
            u_brick_ids = set(
                functools.reduce(set.intersection, data_model_brick_ids.values())
            )
            print(f"# of unique bricks in intersection: {len(u_brick_ids)}")

            # construct DAG run template
            dag_run_template = self.generate_dag_run(
                pmapname=pmapname,
                extra_kwargs=extra_kwargs,
                **dag_group_run,
                **{
                    key: value
                    for key, value in {
                        # "data_model": data_model,
                        "dataset": dataset,
                        "instrument": instrument,
                        "obs_type": obs_type,
                        "obs_group": obs_group,
                        "obs_id": obs_id,
                        "proposal_id": proposal_id,
                        "detector": detector,
                        "filter": filter,
                        "prc_status": prc_status,
                        "qc_status": qc_status,
                        # "custom_id": custom_id,
                        # "healpix": healpix,
                    }.items()
                    if value is not None and value != ""  # 过滤条件
                },
            )
            # construct DAG run list
            dag_run_list = []
            for this_brick_id in u_brick_ids:
                this_dag_run = dag_run_template.copy()
                this_dag_run["custom_id"] = this_brick_id
                if debug:
                    print(this_dag_run)
                dag_run_list.append(this_dag_run)
        else:
            # calculate bricks with required_dags
            # TODO: not implemented yet
            # csst_fs.query_task_state(
            #     filter={
            #         "dag": self.dag,
            #         "custom_id": list(u_brick_ids),
            #     }
            # )
            raise NotImplementedError("Not implemented yet.")
+19 −0
Original line number Diff line number Diff line
@@ -512,6 +512,25 @@ class Dispatcher:
            )
        return task_list

    def dispatch_tile(
        self,
        plan_basis: table.Table,
        data_basis: table.Table,
    ) -> list[dict]:
        """Dispatch tile-level tasks."""
        # return self.dispatch_obsgroup(plan_basis, data_basis)
        pass

    # data_model 依赖
    # 前置程序依赖

    data_model_list = [
        "csst-msc-l2-mbi-mosaic",
    ]
    dag_list = [
        "csst-msc-l2-mbi-mosaic",
    ]

    @staticmethod
    def load_test_data() -> tuple:
        import joblib