"""Modal's model configuration file
Separate into two envs: dev and prod
For dev, update FE options at: https://api-staging.suno.ai/margu/bots/modeltype/

In case a new model is added:
- Add the model to the MODEL_DEV_FNS first
- Deploy the model worker
- Deploy the orchestrator to pick up the worker changes

The keys are model names that will show up in Database.
The values are tuples of (modal function name, modal function call).
The modal function name can be fixed for easy redeploys.
As the actual model changes, we just need to change the key (model name).

Note that this way of deploy has a tiny overlap
since two models are running under the hood shortly in the continuous deploy ~ 10 mins
so that the experiment may not be 100% clean at the start.
So the best practice is to turn off the exp for a bit in conductor
Let the exp warm up and then turn it on
"""

import datetime

ALL_AUDIO_PROMPTS = ["artist", "cover", "overpaint", "underpaint", "future", "history"]

# This is the default mapping from the studio API model names to the engine worker names
DEFAULT_API_TO_MODEL_MAP = {
    # Key is the name that front end will use
    # Value is a dictionary of (cumulative fraction, [model name in db, model name in db])
    # v2 3b model
    "chirp-v2-xxl-alpha": {
        1.0: ["chirp-v2-xxl-alpha", "chirp-v2-xxl-alpha"],
    },
    # v3 7b model
    "chirp-v3-0": {
        1.0: ["chirp-v3-engine-i", "chirp-v3-engine-i"],  # single model
    },
    # v3.5 13b model
    "chirp-v3-5": {
        1.0: ["chirp-v3p5-engine-s-8", "chirp-v3p5-engine-s-8"],
    },
    # v3.5 13b model for audio upload only
    "chirp-v3-5-upload": {
        1.0: ["chirp-v3p5-engine-upload-4", "chirp-v3p5-engine-upload-4"],
    },
    # v3.5 13b model for short (image / video to song)
    "chirp-v3-5-short": {
        1.0: ["chirp-v3p5-engine-short", "chirp-v3p5-engine-short"],
    },
    # v3.5 30b model (supports cover/artist_consistency/infill)
    "chirp-v3-5-tau": {
        1.0: ["chirp-v3p5-engine-t-6", "chirp-v3p5-engine-t-6"],
    },
    # v3.5 13b 2h (for bot traffic)
    "chirp-v3-5-b": {
        1.0: ["chirp-v3p5-engine-b", "chirp-v3p5-engine-b"],
    },
    # v4 diffusion upsample model
    "chirp-up": {
        1.0: ["chirp-v4-up-u-7", "chirp-v4-up-u-7"],
    },
    # v4 13b model
    "chirp-v4": {
        1.0: ["chirp-v4-h-s-32", "chirp-v4-h-s-32"],
    },
    # v4 30b model (supports cover/artist_consistency/infill)
    "chirp-v4-tau": {
        1.0: ["chirp-v4-h-t-6", "chirp-v4-h-t-6"],
    },
    # v4.5 diffusion stem model
    "chirp-stem": {
        1.0: ["chirp-ahi-stem-12-t1", "chirp-ahi-stem-12-t1"],
    },
    # v4.5 diffusion stem complement model
    "chirp-stem-comp": {
        1.0: ["chirp-ahi-stem-comp-t1", "chirp-ahi-stem-comp-t1"],
    },
    # v4? semantic only model
    "chirp-auk": {
        1.0: ["chirp-auk-t1", "chirp-auk-t1"],
    },
    "chirp-auk-infill": {
        1.0: ["chirp-auk-infill", "chirp-auk-infill"],
    },
    "chirp-auk-o": {
        1.0: ["chirp-auk-t0", "chirp-auk-t0"],
    },
    "chirp-auk-test": {
        1.0: ["chirp-auk-dpo", "chirp-auk-dpo"],
    },
    "chirp-seeds": {
        1.0: ["chirp-seeds-v0", "chirp-seeds-v0"],
    },
    # v4.5 diffusion v2 models
    "chirp-ahi": {
        1.0: ["chirp-ahi-up-2", "chirp-ahi-up-2"],
    },
    "chirp-bluejay": {
        1.0: ["chirp-bluejay-t0", "chirp-bluejay-t0"],
    },
}
# this sets up the experimental information for now
# TODO: add more to this and move this to a database
EXPERIMENTAL_INFO = {
    # increase the version number if you update the fractions
    # TODO: no obvious way to auto-increment this nicely
    "experiment_version": "v_931",
    "timestamp": datetime.datetime.now().strftime("%Y%m%d_%H%M%S"),
    # description of the current experiment -- if you update the fractions please update this
    # when you add new experiments, choose the model name carefully
    # fit should be consistent with the base model name, plus some suffix
    # like chirp-v3p5-engine-s-8-xxx, this will help logging and monitoring
    "description": """ 
        testing:
        -
        inference scanning:
        - 13b upload-4
        - 30b tau
        - v4 upsample
        - v4 13b s-32
        - cfg null changes test
        fixing autocast warning
    """,
    # the mapping from the studio API model names to the experiment model names
    "api_model_to_experiment_map": {
        **DEFAULT_API_TO_MODEL_MAP,
        "chirp-v3-5": {
            # 0.025: ["chirp-v3p5-engine-s-31", "chirp-v3p5-engine-s-31-2-bt4"],  # test
            # 0.05: ["chirp-v3p5-engine-s-31", "chirp-v3p5-engine-s-31"],  # data collection
            # 0.10: ["chirp-v3p5-engine-t-6", "chirp-v3p5-engine-t-6"],  # data collection
            # 0.105: ["chirp-v3p5-h-s-31", "chirp-v3p5-h-s-31"],  # test v4 traffic under v3.5
            # 0.15: ["chirp-v3p5-engine-s-8", "chirp-v3p5-engine-s-8-tech-vt-3"],  # test cfg ramp
            # 0.65: ["chirp-v3p5-engine-s-8", "chirp-v3p5-engine-s-8-evict"],  # test evict
            1.0: ["chirp-v3p5-engine-s-8", "chirp-v3p5-engine-s-8"],  # this is good now
        },
        "chirp-v3-5-tau": {
            # 0.33: ["chirp-v3p5-engine-t-6", "chirp-v3p5-engine-t-6-7-2"],
            1.0: ["chirp-v3p5-engine-t-6", "chirp-v3p5-engine-t-6"],
        },
        "chirp-up": {
            1.0: [
                "chirp-v4-up-u-d-2-4",
                "chirp-v4-up-u-d-2-4",
            ],  # for diff v2 data collection -- DO NOT TOUCH yet
            # 1.0: ["chirp-v4-up-u-7", "chirp-v4-up-u-7"],
        },
        "chirp-v4": {
            # 0.15: ["chirp-v4-h-s-32", "chirp-v4-h-s-32-tech-rw-flash3"],
            # 0.025: ["chirp-v4-h-s-32", "chirp-v4-6b-t-01"],  # ab test 30b gen
            # 0.05: ["chirp-v4-h-s-32", "chirp-v4-h-s-32-tech-bct"],  # ab test 30b gen
            # 0.05: ["chirp-v4-6b-t-28", "chirp-v4-6b-t-21"],  # ab test 30b gen
            # 0.05: ["chirp-v4-6b-t-28", "chirp-v4-6b-t-28"],  # ab test 30b gen
            # 0.10: ["chirp-v4-h-s-32", "chirp-v4-h-s-33"],  # ab test new gpt
            # 0.125: ["chirp-v4-h-s-32", "chirp-v4-h-s-32-d-5s"],  # ab test new gpt
            # 0.25: ["chirp-v4-h-s-32", "chirp-v4-h-s-32-tech-zs-1"],  # ab test prev diff
            # 0.125: ["chirp-v4-h-s-32", "chirp-v4-h-s-32-d-10s"],
            # 0.10: ["chirp-v4-h-s-32", "chirp-v4-h-s-33-tech-cn-1"],
            1.0: ["chirp-v4-h-s-32", "chirp-v4-h-s-32"],
        },
        "chirp-v4-tau": {
            # 0.15: ["chirp-v4-h-t-6", "chirp-v4-h-t-6-u-3"],  # ab test prev diff
            # 0.10: ["chirp-v4-6b-t-28", "chirp-v4-6b-t-21"],  # ab test 6b
            # 0.10: ["chirp-v4-6b-t-28", "chirp-v4-6b-t-28"],  # ab test 6b
            # 0.50: [
            #     "chirp-v4-h-t-6",
            #     "chirp-v4-h-t-6-c-35",
            # ],  # THIS IS VERY AGRESSIVE BUT THIS IS ONLY INFILL so fine
            # 0.2: ["chirp-v4-h-t-6", "chirp-v4-h-t-6-cfg-null2"],
            # 0.10: ["chirp-v4-6b-t-21", "chirp-v4-6b-t-21"],  # CFG experiments
            1.0: ["chirp-v4-h-t-6", "chirp-v4-h-t-6"],
        },
        "chirp-auk": {
            0.05: ["chirp-auk-t0", "chirp-auk-t0"],  # for data collection, ask before changing
            0.10: ["chirp-auk-t1", "chirp-auk-t1-d67"],  # ab test auk
            0.15: ["chirp-auk-t1-d7", "chirp-auk-t1-d67"],  # ab test auk
            # 0.20: ["chirp-auk-t1", "chirp-auk-t1-d7"],  # ab test auk
            # 0.25: ["chirp-auk-t1", "chirp-auk-t1-tech-3"],  # ab test modal test
            1.0: ["chirp-auk-t1", "chirp-auk-t1"],
        },
        "chirp-ahi": {
            # 0.5: ["chirp-ahi-up-2", "chirp-ahi-up-d-4-28"],  # data collection
            1.0: ["chirp-ahi-up-2", "chirp-ahi-up-2"],  # data collection
        },
        "chirp-auk-infill": {
            0.5: ["chirp-auk-infill", "chirp-auk-t1-d67"],  # ab test infill
            # 1.0: ["chirp-auk-t1-d7", "chirp-auk-t1-d67"],  # ab test infill
            1.0: ["chirp-auk-infill", "chirp-auk-infill"],
        },
    },
    # the mapping from the studio API model names to model names for bot (API) users
    # typically don't need to change this
    "api_model_to_bot_model_map": {
        **DEFAULT_API_TO_MODEL_MAP,
        "chirp-v3-5": {
            1.0: ["chirp-v3p5-engine-b", "chirp-v3p5-engine-b"],
        },
        # this is for audio upload extention
        "chirp-v3-5-upload": {
            1.0: ["chirp-v3p5-engine-b", "chirp-v3p5-engine-b"],
        },
        # TODO: do something about short audio for bots
        "chirp-v3-5-tau": {
            1.0: ["chirp-v3p5-engine-b-tau", "chirp-v3p5-engine-b-tau"],
        },
        "chirp-v4": {
            1.0: ["chirp-v4-engine-b", "chirp-v4-engine-b"],
        },
        "chirp-v4-tau": {
            1.0: ["chirp-v4-engine-b-tau", "chirp-v4-engine-b-tau"],
        },
        "chirp-up": {
            1.0: ["chirp-v4-up-u-7", "chirp-v4-up-u-7"],
        },
        "chirp-auk": {
            1.0: ["chirp-auk-engine-b", "chirp-auk-engine-b"],
        },
        # TODO: maybe add sth for auk
    },
    # key is original request model name
    # subkey is the experimental model name
    # subsub key is the exp split fraction (cumulative):
    # - logged experiment name for inference
    # - two dict of inference config:
    #   - one for the original model that's empty
    #   - one for the experimental model that's NOT empty
    "inference_experiment": {
        "chirp-v3-5": {
            #  this is a good A/A test stream -- should be the same as prod
            "chirp-v3p5-engine-s-8": {
                0.25: (
                    "min_p_005",  # param experiment name
                    [{"min_p_semantic": 0.005, "min_p_coarse": 0.005}, {}],
                ),
            },
        },
        "chirp-v4": {
            "chirp-v4-h-s-32": {
                # gpt config
                0.05: (
                    "temp_s_70",  # param experiment name
                    [{"temp_semantic": 0.7}, {}],
                ),
                0.10: (
                    "temp_s_80",  # param experiment name
                    [{"temp_semantic": 0.8}, {}],
                ),
                0.15: (
                    "n_tag_3",  # param experiment name
                    [{"n_repeat_tags": 3, "n_repeat_neg_tags": 3}, {}],
                ),
                # diffusion config
                0.20: (
                    "text_1",  # param experiment name
                    [{"text_cfg_coef": 1.0}, {}],
                ),
            },
        },
        "chirp-v4-tau": {
            # Note we can't scan min_p, cause 30b is using TP and torch sampling with top_p
            "chirp-v4-h-t-6": {
                # gpt config
                0.05: (
                    "temp_s_70",  # param experiment name
                    [{"temp_semantic": 0.7}, {}],
                ),
                0.10: (
                    "temp_s_80",  # param experiment name
                    [{"temp_semantic": 0.8}, {}],
                ),
                # # diffusion config
                0.15: (
                    "text_1",  # param experiment name
                    [{"text_cfg_coef": 1.0}, {}],
                ),
            },
        },
        "chirp-up": {
            "chirp-v4-up-u-d-2-4": {
                0.10: (
                    "text_1",  # param experiment name
                    [{"text_cfg_coef": 1.0}, {}],
                ),
                0.20: (
                    "text_3",  # param experiment name
                    [{"text_cfg_coef": 3.0}, {}],
                ),
                0.30: (
                    "step_8",  # param experiment name
                    [{"steps": 8}, {}],
                ),
                0.40: (
                    "step_12",  # param experiment name
                    [{"steps": 12}, {}],
                ),
                0.50: (
                    "noise_l_5",  # param experiment name
                    [{"noise_ctx_level": 0.5}, {}],
                ),
            },
        },
        # 4.5 upsampple scans
        "chirp-ahi": {
            "chirp-ahi-up-2": {
                0.10: (
                    "text_1",  # param experiment name
                    [{"text_cfg_coef": 1.0}, {}],
                ),
                0.20: (
                    "text_4",  # param experiment name
                    [{"text_cfg_coef": 4.0}, {}],
                ),
                0.30: (
                    "step_8",  # param experiment name
                    [{"steps": 8}, {}],
                ),
                0.40: (
                    "step_12",  # param experiment name
                    [{"steps": 12}, {}],
                ),
                0.50: (
                    "step_16",  # param experiment name
                    [{"steps": 16}, {}],
                ),
                0.60: (
                    "noise_l_5",  # param experiment name
                    [{"noise_ctx_level": 0.5}, {}],
                ),
                0.80: (
                    "freedom_55",  # param experiment name
                    [{"semantic_mask_ratio": 0.55}, {}],
                ),
                1.0: (
                    "freedom_75",  # param experiment name
                    [{"semantic_mask_ratio": 0.75}, {}],
                ),
            },
        },
        "chirp-auk": {
            "chirp-auk-t1": {
                # default temp is 0.90
                0.05: (
                    "temp_s_80",  # param experiment name
                    [{"temp_semantic": 0.80}, {}],
                ),
                0.10: (
                    "temp_s_85",  # param experiment name
                    [{"temp_semantic": 0.85}, {}],
                ),
                0.15: (
                    "temp_s_95",  # param experiment name
                    [{"temp_semantic": 0.95}, {}],
                ),
                # default n_repeat_tags is 3
                0.20: (
                    "n_tag_1",  # param experiment name
                    [{"n_repeat_tags": 1, "n_repeat_neg_tags": 1}, {}],
                ),
                0.25: (
                    "n_tag_2",  # param experiment name
                    [{"n_repeat_tags": 2, "n_repeat_neg_tags": 2}, {}],
                ),
                # default tag cfg is 2.0
                0.30: (
                    "tag_cfg_1",  # param experiment name
                    [{"cfg_coef_tags": 1.0}, {}],
                ),
                0.35: (
                    "tag_cfg_3",  # param experiment name
                    [{"cfg_coef_tags": 3.0}, {}],
                ),
                # default cfg steps is 25 * 60 * 2 for gens
                # default cfg steps is 25 * 30 for artist/cover
                0.40: (
                    "cfg_steps_10",  # param experiment name
                    [{"cfg_coef_tags_max_steps": 25 * 10}, {}],
                ),
                0.45: (
                    "cfg_steps_60",  # param experiment name
                    [{"cfg_coef_tags_max_steps": 25 * 60}, {}],
                ),
                0.50: (
                    "cfg_steps_240",  # param experiment name
                    [{"cfg_coef_tags_max_steps": 25 * 60 * 4}, {}],
                ),
            }
        },
    },
}

# These dicts are set as:
# keys: model name
# values: tuple of (modal function name, modal function call)
MODEL_DEV_FNS = {
    "chirp-v2-xxl-alpha": (  # engine v2
        "engine-chirpv2_engine_v2_dev_80s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v3-engine-d": (  # engine v3, dpoed of 3b data
        "engine-chirpv2_engine_7b_dpo_dev_120s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v3-engine-i": (  # 7b ft model ipoed 7b data -- official v3 candidate
        "engine-chirpv2_engine_7b_ipo_dev_120s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v2-engine-msft-60s": (  # engine v3, for msft dev testing
        "engine-chirpv2_engine_7b_ipo_msft_120s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v3p5-engine": (  # 13 raw
        "engine-chirpv2_engine_13b_dev_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v3p5-engine-b": (  # This is bot but normal model to chill them
        "engine-chirpv2_engine_13b_special_8_dev_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v3p5-engine-b-tau": (  # streamed diffusion upsample
        "engine-chirpv2_engine_30b_t6_dev_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-engine-b": (  # This is bot but normal model to chill them
        "engine-chirpv4_engine_13b_special_32_dev_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-engine-b-tau": (  # This is bot but normal model to chill them
        "engine-chirpv4_engine_30b_t6_dev_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v3p5-engine-s-8": (  # 13 2nd dpo model
        "engine-chirpv2_engine_13b_special_8_dev_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v3p5-engine-s-8-tech-vt-3": (  # 13 2nd dpo model
        "engine-chirpv2_engine_13b_tech_dev_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v3p5-engine-short": (  # 13 2nd dpo model
        "engine-chirpv2_engine_13b_short_dev_30s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v3p5-engine-s-31-2-bt4": (  # 13 3rd dpo model
        "engine-chirpv2_engine_13b_special_test_dev_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v3p5-engine-upload-4": (  # 13b extend model
        "engine-chirpv2_engine_13b_upload_4_dev_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v3p5-engine-t-1": (  # 30 raw -- this is v4 but we hide it
        "engine-chirpv2_engine_30b_dev_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v3p5-engine-t-c": (  # 30 classical
        "engine-chirpv2_engine_30b_classical_dev_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v3p5-engine-t-dance": (  # 30 dance
        "engine-chirpv2_engine_30b_dance_dev_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v3p5-engine-s-31": (  # 13 3rd dpo model
        "engine-chirpv2_engine_13b_special_31_dev_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v3p5-engine-t-6": (  # streamed diffusion upsample
        "engine-chirpv2_engine_30b_t6_dev_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-up-u-7": (  # upsample model
        "upsample-diff_v1-dev",
        "UpsampleStub.upsample",
    ),
    "chirp-v4-up-u-d-2-4": (  # upsample model, masked as v4, data-2
        "upsample-diff_v2_data-dev",
        "UpsampleStub.upsample",
    ),
    "chirp-ahi-up-d-4-28": (  # upsample model, masked as v4, data-2
        "upsample-diff_v2_test-dev",
        "UpsampleStub.upsample",
    ),
    "chirp-ahi-up-2": (  # upsample model, masked as v4, data-2
        "upsample-diff_v2_data-dev",
        "UpsampleStub.upsample",
    ),
    "chirp-seeds-v0": (
        "upsample-diff_seeds_v0-dev",
        "UpsampleStub.upsample",
    ),
    "chirp-v3p5-h-s-31": (  # streamed diffusion upsample
        "engine-chirpv4_engine_13b_special_31_dev_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v3p5-h-t-6": (  # streamed diffusion upsample
        "engine-chirpv4_engine_30b_t6_dev_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-h-s-31": (  # streamed diffusion upsample
        "engine-chirpv4_engine_13b_special_31_dev_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-h-s-32": (  # streamed diffusion upsample
        "engine-chirpv4_engine_13b_special_32_dev_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-h-s-33": (  # streamed diffusion upsample
        "engine-chirpv4_engine_13b_special_33_dev_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-h-s-32-d-5s": (  # streamed diffusion upsample
        "engine-chirpv4_engine_13b_special_test_dev_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-h-s-32-u-6-3": (  # streamed diffusion upsample
        "engine-chirpv4_engine_13b_special_test_diff_dev_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-h-s-32-tech-zs-1": (
        "engine-chirpv4_engine_13b_tech_dev_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-h-s-32-d-10s": (  # 13 3rd dpo model, test diff
        "engine-chirpv4_engine_13b_special_test_2_dev_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-h-s-32-c-30-p-1-1-f": (  # 13 3rd dpo model, test diff
        "engine-chirpv4_engine_13b_special_test_3_dev_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-h-t-6-c-35": (  # streamed diffusion upsample
        "engine-chirpv4_engine_30b_t6_test_dev_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-h-t-6-u-3": (  # streamed diffusion upsample
        "engine-chirpv4_engine_30b_test_diff_dev_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-h-t-6": (  # streamed diffusion upsample
        "engine-chirpv4_engine_30b_t6_dev_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-h-t-6-tech-gk-2": (  # streamed diffusion upsample
        "engine-chirpv4_engine_30b_t6_tech_dev_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-h-t-6-tech-tune": (  # streamed diffusion upsample
        "engine-chirpv4_engine_30b_t6_tech_dev_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-h-t-6-cfg-null2": (  # streamed diffusion upsample
        "engine-chirpv4_engine_30b_t6_tech_dev_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-h-t-6-eval": (
        "engine-chirpv4_engine_30b_t6_eval_dev_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-h-api": (  # streamed diffusion upsample
        "engine-chirpv4_engine_13b_special_32_fast_dev_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-ahi-stem-comp-t1": (
        "stems-stems_v1-dev",
        "StemStub.stem",
    ),
    "chirp-ahi-stem-12-t1": (
        "stems-stems_v1_12_output-dev",
        "StemStub.stem",
    ),
    "chirp-auk-orig": (
        "engine-chirpv4_engine_6b_sem_dev_480s",
        "ChirpV2Stub.generate",
    ),
    "chirp-auk-t0": (
        "engine-chirpv4_engine_6b_sem_dev_480s",
        "ChirpV2Stub.generate",
    ),
    "chirp-bluejay-t0": (
        "engine-chirpv4_engine_6b_sem_t2_dev_480s",
        "ChirpV2Stub.generate",
    ),
    "chirp-bluejay-a": (  # bluejay dpo t1-r7
        "engine-chirpv4_engine_6b_sem_bluejay_test_dev_480s",
        "ChirpV2Stub.generate",
    ),
    "chirp-bluejay-c": (  # bluejay sft
        "engine-chirpv4_engine_6b_sem_bluejay_dev_480s",
        "ChirpV2Stub.generate",
    ),
    "chirp-bluejay-d": (  # auk 3b
        "engine-chirpv4_engine_3b_sem_orig_dev_480s",
        "ChirpV2Stub.generate",
    ),
    "chirp-bluejay-e": (  # bluejay dpo t0-v6
        "engine-chirpv4_engine_6b_sem_bluejay_test_2_dev_480s",
        "ChirpV2Stub.generate",
    ),
    "chirp-bluejay-f": (  # bluejay dpo t0-v7
        "engine-chirpv4_engine_6b_sem_bluejay_test_2_dev_480s",
        "ChirpV2Stub.generate",
    ),
    "chirp-bluejay-g": (  # bluejay t1-r9-d28
        "engine-chirpv4_engine_6b_sem_bluejay_dev_480s",
        "ChirpV2Stub.generate",
    ),
    "chirp-bluejay-h": (  # bluejay dpo d60
        "engine-chirpv4_engine_6b_sem_bluejay_test_2_dev_480s",
        "ChirpV2Stub.generate",
    ),
    "chirp-auk-task-test": (
        "engine-chirpv4_engine_6b_sem_task_dev_480s",
        "ChirpV2Stub.generate",
    ),
    "chirp-auk-eval": (
        "engine-chirpv4_engine_6b_sem_eval_dev_480s",
        "ChirpV2Stub.generate",
    ),
    "chirp-auk-dpo": (
        "engine-chirpv4_engine_6b_sem_test_dev_480s",
        "ChirpV2Stub.generate",
    ),
    "chirp-auk-t1-d7": (
        "engine-chirpv4_engine_6b_sem_test_2_dev_480s",
        "ChirpV2Stub.generate",
    ),
    "chirp-auk-t1-d67": (
        "engine-chirpv4_engine_6b_sem_test_dev_480s",
        "ChirpV2Stub.generate",
    ),
    "chirp-auk-engine-b": (
        "engine-chirpv4_engine_6b_sem_t1_dev_480s",
        "ChirpV2Stub.generate",
    ),
    "chirp-auk-t1": (
        "engine-chirpv4_engine_6b_sem_t1_dev_480s",
        "ChirpV2Stub.generate",
    ),
    "chirp-auk-infill": (
        "engine-chirpv4_engine_30b_t6_infill_dev_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-auk-t1-d58": (  # auk 6b tech
        "engine-chirpv4_engine_6b_sem_tech_dev_480s",
        "ChirpV2Stub.generate",
    ),
    "chirp-auk-t1-tech-3": (  # auk 6b tech modal
        "engine-chirpv4_engine_6b_sem_tech_dev_480s",
        "ChirpV2Stub.generate",
    ),
}
# These are the prod envs
MODEL_PROD_FNS = {
    "chirp-v2-engine-msft-60s": (  # engine v3
        # change this modal endpoint path if we need to update
        "engine-chirpv2_engine_7b_ipo_msft_120s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v2-xxl-alpha": (  # engine v2
        "engine-chirpv2_engine_v2_prod_80s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v3-engine-i": (  # 7b ft ipoed 7b data, v3.0
        "engine-chirpv2_engine_7b_ipo_prod_120s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v3p5-engine-s-8": (  # 13 dpo comparison model
        "engine-chirpv2_engine_13b_special_8_prod_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v3p5-engine-s-8-tech-vt-3": (  # 13 2nd dpo model
        "engine-chirpv2_engine_13b_tech_prod_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v3p5-engine-short": (  # 13 dpo comparison model
        "engine-chirpv2_engine_13b_short_prod_30s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v3p5-engine-s-31-2-bt4": (  # 13 3rd dpo model
        "engine-chirpv2_engine_13b_special_test_prod_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-h-s-32-u-6-3": (  # streamed diffusion upsample
        "engine-chirpv4_engine_13b_special_test_diff_prod_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v3p5-engine-upload-4": (  # 13b extend model
        "engine-chirpv2_engine_13b_upload_4_prod_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v3p5-engine-b": (  # This is bot but normal model to chill them
        "engine-chirpv2_engine_13b_special_8_prod_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v3p5-engine-b-tau": (  # This is bot but normal model to chill them
        "engine-chirpv2_engine_30b_t6_prod_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-engine-b": (  # This is bot but normal model to chill them
        "engine-chirpv4_engine_13b_special_32_prod_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-engine-b-tau": (  # This is bot but normal model to chill them
        "engine-chirpv4_engine_30b_t6_prod_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v3p5-engine-s-31": (  # 13 3rd dpo model
        "engine-chirpv2_engine_13b_special_31_prod_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v3p5-engine-t-6": (  # streamed diffusion upsample
        "engine-chirpv2_engine_30b_t6_prod_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v3p5-engine-t-6-7-2": (  # 30 dpo test
        "engine-chirpv2_engine_30b_t6_test_prod_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-up-u-7": (  # upsample model
        "upsample-diff_v1-prod",
        "UpsampleStub.upsample",
    ),
    "chirp-v4-up-u-d-2-4": (  # upsample model, masked as v4, data-2
        "upsample-diff_v2_data-prod",
        "UpsampleStub.upsample",
    ),
    "chirp-ahi-up-d-4-28": (  # upsample model, masked as v4, data-2
        "upsample-diff_v2_test-prod",
        "UpsampleStub.upsample",
    ),
    "chirp-ahi-up-2": (  # upsample model, masked as v4, data-2
        "upsample-diff_v2_data-prod",
        "UpsampleStub.upsample",
    ),
    "chirp-v3p5-h-s-31": (  # streamed diffusion upsample
        "engine-chirpv4_engine_13b_special_31_prod_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v3p5-h-t-6": (  # streamed diffusion upsample
        "engine-chirpv4_engine_30b_t6_prod_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-h-s-31": (  # streamed diffusion upsample
        "engine-chirpv4_engine_13b_special_31_prod_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-h-s-32": (  # streamed diffusion upsample
        "engine-chirpv4_engine_13b_special_32_prod_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-h-s-33": (  # streamed diffusion upsample
        "engine-chirpv4_engine_13b_special_33_prod_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-h-s-32-d-5s": (  # streamed diffusion upsample
        "engine-chirpv4_engine_13b_special_test_prod_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-h-s-32-d-10s": (  # 13 3rd dpo model, test diff
        "engine-chirpv4_engine_13b_special_test_2_prod_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-h-s-32-tech-rw-flash3": (  # 13 3rd dpo model, test diff
        "engine-chirpv4_engine_13b_special_test_tech_1_prod_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-h-s-32-tech-zs-1": (
        "engine-chirpv4_engine_13b_tech_prod_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-h-s-32-c-30-p-1-1-f": (  # 13 3rd dpo model, test diff
        "engine-chirpv4_engine_13b_special_test_3_prod_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-h-t-6-c-35": (  # streamed diffusion upsample
        "engine-chirpv4_engine_30b_t6_test_prod_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-h-t-6-u-3": (  # streamed diffusion upsample
        "engine-chirpv4_engine_30b_test_diff_prod_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-h-t-6-tech-gk-2": (  # streamed diffusion upsample
        "engine-chirpv4_engine_30b_t6_tech_prod_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-h-t-6-tech-tune": (  # streamed diffusion upsample
        "engine-chirpv4_engine_30b_t6_tech_prod_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-h-t-6-cfg-null2": (  # streamed diffusion upsample
        "engine-chirpv4_engine_30b_t6_tech_prod_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-h-t-6": (  # streamed diffusion upsample
        "engine-chirpv4_engine_30b_t6_prod_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-v4-h-api": (  # streamed diffusion upsample
        "engine-chirpv4_engine_13b_special_32_fast_prod_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-ahi-stem-comp-t1": (
        "stems-stems_v1-prod",
        "StemStub.stem",
    ),
    "chirp-ahi-stem-12-t1": (
        "stems-stems_v1_12_output-prod",
        "StemStub.stem",
    ),
    "chirp-auk-t0": (
        "engine-chirpv4_engine_6b_sem_prod_480s",
        "ChirpV2Stub.generate",
    ),
    "chirp-bluejay-t0": (
        "engine-chirpv4_engine_6b_sem_t2_prod_480s",
        "ChirpV2Stub.generate",
    ),
    "chirp-auk-engine-b": (
        "engine-chirpv4_engine_6b_sem_t1_prod_480s",
        "ChirpV2Stub.generate",
    ),
    "chirp-auk-t1": (
        "engine-chirpv4_engine_6b_sem_t1_prod_480s",
        "ChirpV2Stub.generate",
    ),
    "chirp-auk-t1-d7": (
        "engine-chirpv4_engine_6b_sem_test_2_prod_480s",
        "ChirpV2Stub.generate",
    ),
    "chirp-auk-t1-d67": (
        "engine-chirpv4_engine_6b_sem_test_prod_480s",
        "ChirpV2Stub.generate",
    ),
    "chirp-auk-infill": (
        "engine-chirpv4_engine_30b_t6_infill_prod_240s",
        "ChirpV2Stub.generate",
    ),
    "chirp-auk-t1-d58": (  # auk 6b tech
        "engine-chirpv4_engine_6b_sem_tech_prod_480s",
        "ChirpV2Stub.generate",
    ),
    "chirp-auk-t1-tech-3": (  # auk 6b tech
        "engine-chirpv4_engine_6b_sem_tech_prod_480s",
        "ChirpV2Stub.generate",
    ),
}


def _validate_model_and_experimental_config():
    """Validate the model and experimental config."""
    set_of_models = set()
    list_of_experiment_maps = [
        "api_model_to_experiment_map",
        "api_model_to_bot_model_map",
    ]
    for experiment_map in list_of_experiment_maps:
        for model_name, model_config in EXPERIMENTAL_INFO[experiment_map].items():
            prev_exp_fraction = 0
            for exp_fraction, exp_model_names in model_config.items():
                if exp_fraction < prev_exp_fraction:
                    raise ValueError(
                        f"Experimental fractions are not in order for {model_name} -- {exp_fraction}!"
                    )
                prev_exp_fraction = exp_fraction
                for exp_model_name in exp_model_names:
                    set_of_models.add(exp_model_name)
    for model_name, model_config in EXPERIMENTAL_INFO["inference_experiment"].items():
        for exp_model_name in model_config.keys():
            set_of_models.add(exp_model_name)

    for model_name in set_of_models:
        if model_name not in MODEL_DEV_FNS or model_name not in MODEL_PROD_FNS:
            # exceptions for modle we don't want to have on prod
            if (
                ("chirp-auk" in model_name)
                or ("chirp-auk" in model_name)
                or ("chirp-seeds" in model_name)
            ):
                continue
            raise ValueError(f"Model {model_name} not found in MODEL_DEV_FNS or MODEL_PROD_FNS!")
    for _, (modal_fn, _) in MODEL_DEV_FNS.items():
        if ("dev" not in modal_fn) and ("msft" not in modal_fn):
            raise ValueError(f"Modal {modal_fn} doesn't contain dev!")
    for _, (modal_fn, _) in MODEL_PROD_FNS.items():
        if ("prod" not in modal_fn) and ("msft" not in modal_fn):
            raise ValueError(f"Modal {modal_fn} doesn't contain prod!")


# always excute this function to validate the model and experimental config
_validate_model_and_experimental_config()
