diff --git a/README.md b/README.md index c89945f..1f4b5c7 100644 --- a/README.md +++ b/README.md @@ -105,14 +105,8 @@ or you can get the `run_id` via the MLFlow UI. MLFLOW_TRACKING_URI=http://localhost:8080 uv run modelplane annotate --annotator_id {annotator_id} --experiment expname --response_run_id {run_id} ``` -#### Private Ensemble +#### Private annotator If you have access to the private annotator, you can run directly with: ``` -MLFLOW_TRACKING_URI=http://localhost:8080 uv run modelplane annotate --annotator_id safety-v1.1 --experiment expname --response_run_id {run_id} -``` - - -### Custom Ensembles -``` -MLFLOW_TRACKING_URI=http://localhost:8080 uv run modelplane annotate --annotator_id {annotator_id1} --annotator_id {annotator_id2} --ensemble_strategy {ensemble_strategy} --experiment expname --response_file path/to/response.csv +MLFLOW_TRACKING_URI=http://localhost:8080 uv run modelplane annotate --annotator_id safety-v1.2 --experiment expname --response_run_id {run_id} ``` diff --git a/flightpaths/Ensemble Development Template.ipynb b/flightpaths/Ensemble Development Template.ipynb deleted file mode 100644 index 3060cfc..0000000 --- a/flightpaths/Ensemble Development Template.ipynb +++ /dev/null @@ -1,287 +0,0 @@ -{ - "cells": [ - { - "cell_type": "markdown", - "id": "ab195250-6a0f-4176-a09d-3696d911203d", - "metadata": {}, - "source": [ - "# Ensemble Evaluator Development Template\n", - "\n", - "Use this template to iterate on a new ensemble. The Evaluator Development Template notebook may be helpful as a precursor to understanding the flow here." - ] - }, - { - "cell_type": "markdown", - "id": "3d2d5865-2cd7-4b81-a588-dfec27727643", - "metadata": {}, - "source": [ - "## Imports" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "id": "f44e837c-05e9-4e62-916d-9884bb47839e", - "metadata": {}, - "outputs": [], - "source": [ - "import datetime\n", - "from modelplane.runways import responder, annotator, scorer" - ] - }, - { - "cell_type": "markdown", - "id": "726c8897-db04-4435-8d67-7a05309ef740", - "metadata": {}, - "source": [ - "## Settings\n", - "\n", - "The `sut_id` refers to the model that generates the responses to the prompts. \n", - "\n", - "The `experiment` variable will be used to organize the various runs in mlflow.\n", - "\n", - "The `prompts` should point to a location in `/flightpaths/data`. A sample dataset is provided.\n", - "\n", - "Finally, `num_workers` can adjust the parallelism." - ] - }, - { - "cell_type": "code", - "execution_count": null, - "id": "414f9c85-d146-4119-854b-e009235aa4c4", - "metadata": {}, - "outputs": [], - "source": [ - "sut_id = \"demo_yes_no\"\n", - "experiment = \"ensemble_experiment_\" + datetime.date.today().strftime(\"%Y%m%d\")\n", - "prompts = \"data/airr_official_1.0_demo_en_us_prompt_set_release_reduced.csv\"\n", - "ground_truth = \"data/fakegroundtruth.csv\"\n", - "num_workers = 4" - ] - }, - { - "cell_type": "markdown", - "id": "524e0529-4e51-45ae-b2ab-313915881f98", - "metadata": {}, - "source": [ - "## Defining a new ensemble\n", - "\n", - "An ensemble is a collection of annotators together with combining logic.\n", - "\n", - "* An implementation of an annotator should define:\n", - "\n", - " * `translate_prompt`: takes a `TextPrompt` and `SUTResponse` (which have the needed information on the prompt responses for the given LLM) and translates the request into an arbitrary form useful for the `annotate` method.\n", - " * `annotate`: take the translated request and produce an annotation in any form.\n", - " * `translate_response`: take the annotation and translate it to the needed from (`SafetyModelResponse`).\n", - "\n", - "* An implementation of the combining logic should define `compute_response` which takes a map of `SafetyAnnotation`s from the underlying annotators and combines them to produce a final annotation (`EnsembleSafetyAnnotation`, which is simply a `SafetyAnnotation` along with the map of underlying responses).\n", - "\n", - "Below, we give a sample ensemble, which takes two independent `RandomAnnotator` responses and reports unsafe if any underlying annotation was unsafe. This requires implementations of the `RandomAnnotator` and `AnyUnsafeStrategy`." - ] - }, - { - "cell_type": "markdown", - "id": "dc8644ad-2a9b-4953-ada5-23b252026391", - "metadata": {}, - "source": [ - "### Annotators - `RandomAnnotator`\n", - "\n", - "Of course, if the annotators are already defined, you can access them directly by their annotator ids.\n", - "\n", - "Annotators must be registered using the annotator registry `ANNOTATORS`." - ] - }, - { - "cell_type": "code", - "execution_count": null, - "id": "89674a6d-b2c5-42a3-9a0c-927101126877", - "metadata": {}, - "outputs": [], - "source": [ - "import random\n", - "\n", - "from modelgauge.annotation import SafetyAnnotation\n", - "from modelgauge.annotator import Annotator\n", - "from modelgauge.annotator_registry import ANNOTATORS\n", - "from modelgauge.prompt import ChatPrompt, TextPrompt\n", - "from modelgauge.sut import SUTResponse\n", - "from modelgauge.tests.safe_v1 import BaseSafeTestVersion1\n", - "from pydantic import BaseModel\n", - "\n", - "\n", - "class AnnotationRequest(BaseModel):\n", - " prompt: ChatPrompt | TextPrompt\n", - " response: SUTResponse\n", - "\n", - "\n", - "class RandomAnnotator(Annotator):\n", - " \"\"\"Simple random implementation for demonstration. Uses SafetyModelResponse\n", - " for consistency with AILuminate benchmarks.\n", - "\n", - " The methods to implement are based on the idea that annotate will \n", - " potentially make a call to another system and the translators are helpful \n", - " in preparing that call and converting the response to the needed annotation.\n", - " \"\"\"\n", - " def __init__(self, uid, prob_unsafe):\n", - " super().__init__(uid)\n", - " self.prob_unsafe = prob_unsafe\n", - " \n", - " def translate_prompt(self, prompt: TextPrompt | ChatPrompt, response: SUTResponse):\n", - " return AnnotationRequest(\n", - " prompt=prompt,\n", - " response=response,\n", - " )\n", - "\n", - " def annotate(self, annotation_request: AnnotationRequest) -> SafetyAnnotation:\n", - " unsafe_categories = []\n", - " for category in BaseSafeTestVersion1.hazards:\n", - " if random.random() < self.prob_unsafe:\n", - " unsafe_categories.append(category)\n", - " return SafetyAnnotation(\n", - " is_safe=not unsafe_categories,\n", - " is_valid=True,\n", - " )\n", - "\n", - " def translate_response(self, request: AnnotationRequest, response: SafetyAnnotation) -> SafetyAnnotation:\n", - " return response\n", - "\n", - "\n", - "annotator_id1 = \"random01\"\n", - "ANNOTATORS.register(RandomAnnotator, annotator_id1, 0.1)\n", - "\n", - "annotator_id2 = \"random02\"\n", - "ANNOTATORS.register(RandomAnnotator, annotator_id2, 0.2)" - ] - }, - { - "cell_type": "markdown", - "id": "43bc96aa-7a5c-4ac4-ab2d-ad89c96b4a28", - "metadata": {}, - "source": [ - "### Combiner logic - `AnyUnsafeStrategy`\n", - "\n", - "This implementation will return is_safe = \"unsafe\" if any of the ensemble members decide it's unsafe.\n", - "\n", - "Combiner logic must be declared via `ENSEMBLE_STRATEGIES`." - ] - }, - { - "cell_type": "code", - "execution_count": null, - "id": "a9170875-d822-4e49-8445-37639caee836", - "metadata": {}, - "outputs": [], - "source": [ - "from typing import Dict\n", - "\n", - "from modelgauge.annotation import EnsembleSafetyAnnotation\n", - "from modelgauge.ensemble_strategies import ENSEMBLE_STRATEGIES, EnsembleStrategy\n", - "\n", - "\n", - "class AnyUnsafeEnsembleStrategy(EnsembleStrategy):\n", - " def compute_response(self, annotations: Dict[str, SafetyAnnotation]) -> EnsembleSafetyAnnotation:\n", - " response = EnsembleSafetyAnnotation(\n", - " joined_responses=annotations, \n", - " is_safe=True,\n", - " is_valid=False,\n", - " )\n", - " for annotation in annotations.values():\n", - " if annotation.is_valid:\n", - " response.is_valid = True\n", - " if not annotation.is_safe:\n", - " response.is_safe = False\n", - " return response\n", - "\n", - "ENSEMBLE_STRATEGIES[\"anyunsafe\"] = AnyUnsafeEnsembleStrategy()" - ] - }, - { - "cell_type": "markdown", - "id": "17760cd3-23fe-4c79-8882-475d8d7096ea", - "metadata": {}, - "source": [ - "## Run the model\n", - "\n", - "This step will get responses to the prompts from the given SUT.\n", - "\n", - "Save this run_id to avoid having to re-run the model later. The results are saved as an artifact in mlflow.\n", - "\n", - "You can see the the runs associated with the experiment you specified above at your MLFlow tracking server. There will be a link below after running. You may need to replace http://mlflow:8080 with http://localhost:8080 in the generated mlflow links if you used the default settings." - ] - }, - { - "cell_type": "code", - "execution_count": null, - "id": "b70d76d5-a3e1-4cc0-aeff-e71b6ff64825", - "metadata": {}, - "outputs": [], - "source": [ - "response_run = responder.respond(\n", - " sut_id=sut_id,\n", - " experiment=experiment,\n", - " prompts=prompts,\n", - " num_workers=num_workers,\n", - ")" - ] - }, - { - "cell_type": "markdown", - "id": "740a8a85-c171-4d11-b094-cd617b14b6ed", - "metadata": {}, - "source": [ - "## Annotate the model\n", - "\n", - "This will run the evaluator on the responses from the prior step. You'll be able to see the details of the annotations in mlflow." - ] - }, - { - "cell_type": "code", - "execution_count": null, - "id": "06632c4d-90bd-4c2d-9c36-84e59dd8f190", - "metadata": {}, - "outputs": [], - "source": [ - "annotation_run = annotator.annotate(\n", - " annotator_ids=[annotator_id1, annotator_id2],\n", - " ensemble_strategy=\"anyunsafe\",\n", - " experiment=experiment,\n", - " response_run_id=response_run.run_id,\n", - " num_workers=num_workers,\n", - ")" - ] - }, - { - "cell_type": "markdown", - "id": "555f95d5-95ae-4919-96b0-55f2378d5846", - "metadata": {}, - "source": [ - "## Score the model\n", - "\n", - "Compute metrics against the given ground truth dataset." - ] - }, - { - "cell_type": "code", - "execution_count": null, - "id": "b4b8672c-32b5-45b1-9be2-396ceb857f44", - "metadata": {}, - "outputs": [], - "source": [ - "scorer.score(\n", - " annotation_run_id=annotation_run.run_id,\n", - " experiment=experiment,\n", - " ground_truth=ground_truth,\n", - ")" - ] - } - ], - "metadata": { - "kernelspec": { - "display_name": "Python 3 (ipykernel)", - "language": "python", - "name": "python3" - } - }, - "nbformat": 4, - "nbformat_minor": 5 -} diff --git a/src/modelplane/cli.py b/src/modelplane/cli.py index c08d50a..d802f1d 100644 --- a/src/modelplane/cli.py +++ b/src/modelplane/cli.py @@ -2,12 +2,10 @@ import click from modelgauge.data_schema import AnnotationSchema -from modelgauge.ensemble_strategies import ENSEMBLE_STRATEGIES from modelplane.runways.annotator import annotate from modelplane.runways.lister import ( list_annotators, - list_ensemble_strategies, list_suts, ) from modelplane.runways.responder import respond @@ -27,11 +25,6 @@ def list_annotators_cli(): list_annotators() -@cli.command(name="list-ensemble-strategies", help="List known ensemble strategies.") -def list_ensemble_strategies_cli(): - list_ensemble_strategies() - - @cli.command(name="list-suts", help="List known suts.") def list_suts_cli(): list_suts() @@ -146,13 +139,6 @@ def get_sut_responses( default=None, help="The annotator UID(s) to use. Multiple annotators can be specified.", ) -@click.option( - "--ensemble_strategy", - type=str, - default=None, - help="The ensemble strategy to use. If set, individual annotator results will be combined using the given strategy. " - "Available strategies: " + ", ".join(list(ENSEMBLE_STRATEGIES.keys())), -) @click.option( "--overwrite", is_flag=True, @@ -202,7 +188,6 @@ def get_annotations( dvc_repo: str | None = None, response_file: str | None = None, response_run_id: str | None = None, - ensemble_strategy: str | None = None, overwrite: bool = False, disable_cache: bool = False, num_workers: int = 1, @@ -217,7 +202,6 @@ def get_annotations( response_file=response_file, response_run_id=response_run_id, annotator_ids=annotator_id, - ensemble_strategy=ensemble_strategy, overwrite=overwrite, disable_cache=disable_cache, num_workers=num_workers, diff --git a/src/modelplane/runways/annotator.py b/src/modelplane/runways/annotator.py index 22ff5b6..c0eedb3 100644 --- a/src/modelplane/runways/annotator.py +++ b/src/modelplane/runways/annotator.py @@ -13,8 +13,6 @@ from modelgauge.annotator import Annotator from modelgauge.annotator_registry import ANNOTATORS from modelgauge.dataset import AnnotationDataset -from modelgauge.ensemble_annotator import EnsembleAnnotator -from modelgauge.ensemble_strategies import ENSEMBLE_STRATEGIES from modelgauge.pipeline_runner import build_runner from modelplane.mlflow.loghelpers import log_tags @@ -36,9 +34,6 @@ ) -DEFAULT_ENSEMBLE_ANNOTATOR_UID = "ensemble" - - def annotate( experiment: str, annotator_ids: List[str], @@ -46,7 +41,6 @@ def annotate( dvc_repo: str | None = None, response_file: str | None = None, response_run_id: str | None = None, - ensemble_strategy: str | None = None, overwrite: bool = False, disable_cache: bool = False, num_workers: int = 1, @@ -58,8 +52,7 @@ def annotate( """ Run annotations and record measurements. """ - # this will set annotator_ids and optionally ensemble - pipeline_kwargs = _get_annotator_settings(annotator_ids, ensemble_strategy) + pipeline_kwargs = _get_annotator_settings(annotator_ids) if not disable_cache: pipeline_kwargs["cache_dir"] = CACHE_DIR pipeline_kwargs["num_workers"] = num_workers @@ -73,8 +66,6 @@ def annotate( for annotator_id in pipeline_kwargs["annotators"] } ) - if ensemble_strategy is not None: - tags["ensemble_strategy"] = ensemble_strategy experiment_id = get_experiment_id(experiment) if overwrite and response_run_id: @@ -124,11 +115,7 @@ def annotate( # log summary statistics annotator_uids = sorted(pipeline_kwargs["annotators"].keys()) log_safety_summary( - annotator_uids=( - annotator_uids - if ensemble_strategy is None - else annotator_uids + [DEFAULT_ENSEMBLE_ANNOTATOR_UID] - ), + annotator_uids=annotator_uids, data_path=pipeline_runner.output_dir() / pipeline_runner.output_file_name, dir=tmp, @@ -144,28 +131,8 @@ def annotate( return RunArtifacts(run_id=run.info.run_id, artifacts=artifacts) -def _get_annotator_settings( - annotator_ids: List[str], - ensemble_strategy: str | None, -) -> Dict[str, Any]: - - kwargs = {} - - kwargs["annotators"] = _get_annotators(annotator_ids) - - if ensemble_strategy is not None: - if ensemble_strategy not in ENSEMBLE_STRATEGIES: - raise ValueError( - f"Unknown ensemble strategy: {ensemble_strategy}. " - f"Available strategies: {list(ENSEMBLE_STRATEGIES.keys())}" - ) - annotator = EnsembleAnnotator( - uid=DEFAULT_ENSEMBLE_ANNOTATOR_UID, - annotators=annotator_ids, - ensemble_strategy=ensemble_strategy, - ) - kwargs["annotators"][DEFAULT_ENSEMBLE_ANNOTATOR_UID] = annotator - return kwargs +def _get_annotator_settings(annotator_ids: List[str]) -> Dict[str, Any]: + return {"annotators": _get_annotators(annotator_ids)} def _get_annotators(annotator_ids: List[str]) -> Dict[str, Annotator]: @@ -202,7 +169,7 @@ def log_safety_summary( mlflow.log_metric(f"{annotator_uid}_total_safe", total_safe[annotator_uid]) mlflow.log_metric(f"{annotator_uid}_total_count", total[annotator_uid]) # TODO: the format for the log probs isn't always the same - # in particular, the private ensemble uses a different format + # in particular, the private annotator uses a different format try: log_stats(f"{annotator_uid}_logprobs_", all_log_probs[annotator_uid]) log_hist(dir, f"{annotator_uid}", all_log_probs[annotator_uid]) diff --git a/src/modelplane/runways/lister.py b/src/modelplane/runways/lister.py index b28b517..a713f7a 100644 --- a/src/modelplane/runways/lister.py +++ b/src/modelplane/runways/lister.py @@ -1,5 +1,4 @@ from modelgauge.annotator_registry import ANNOTATORS -from modelgauge.ensemble_strategies import ENSEMBLE_STRATEGIES from modelgauge.sut_registry import SUTS @@ -9,7 +8,3 @@ def list_annotators(): def list_suts(): print(SUTS.compact_uid_list()) - - -def list_ensemble_strategies(): - print(sorted(ENSEMBLE_STRATEGIES)) diff --git a/tests/it/test_cli.py b/tests/it/test_cli.py index 55f9205..4e8c427 100644 --- a/tests/it/test_cli.py +++ b/tests/it/test_cli.py @@ -25,7 +25,6 @@ def test_main_help(): "score", "list-suts", "list-annotators", - "list-ensemble-strategies", ], ) def test_command_help(command): diff --git a/tests/unit/test_lister.py b/tests/unit/test_lister.py index b879a47..fd91b98 100644 --- a/tests/unit/test_lister.py +++ b/tests/unit/test_lister.py @@ -1,8 +1,5 @@ -from modelgauge.ensemble_strategies import ENSEMBLE_STRATEGIES - from modelplane.runways.lister import ( list_annotators, - list_ensemble_strategies, list_suts, ) @@ -13,15 +10,6 @@ def test_list_annotators(capsys): assert "demo_annotator" in output -def test_list_ensemble_strategies(capsys): - ENSEMBLE_STRATEGIES["demo_ensemble_strategy"] = ENSEMBLE_STRATEGIES["any_unsafe"] - list_ensemble_strategies() - output = capsys.readouterr().out.strip() - assert "demo_ensemble_strategy" in output - - del ENSEMBLE_STRATEGIES["demo_ensemble_strategy"] # Clean up after test - - def test_list_suts(capsys): list_suts() output = capsys.readouterr().out.strip()