{
 "cells": [
  {
   "cell_type": "markdown",
   "id": "0",
   "metadata": {},
   "source": [
    "# Generating GCG Suffixes Using Azure Machine Learning"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "1",
   "metadata": {},
   "source": [
    "This notebook shows how to generate GCG [@zou2023gcg] suffixes using Azure Machine Learning (AML), which consists of three main steps:\n",
    "1. Connect to an Azure Machine Learning (AML) workspace.\n",
    "2. Create AML Environment with the Python dependencies.\n",
    "3. Submit a training job to AML."
   ]
  },
  {
   "cell_type": "markdown",
   "id": "2",
   "metadata": {},
   "source": [
    "## Connect to Azure Machine Learning Workspace"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "3",
   "metadata": {},
   "source": [
    "The [workspace](https://docs.microsoft.com/en-us/azure/machine-learning/concept-workspace) is the top-level resource for Azure Machine Learning (AML), providing a centralized place to work with all the artifacts you create when using AML. In this section, we will connect to the workspace in which the job will be run.\n",
    "\n",
    "To connect to a workspace, we need identifier parameters - a subscription, resource group and workspace name. We will use these details in the `MLClient` from `azure.ai.ml` to get a handle to the required AML workspace. We use the [default Azure authentication](https://docs.microsoft.com/en-us/python/api/azure-identity/azure.identity.defaultazurecredential?view=azure-python) for this tutorial."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "4",
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "Found default environment files: ['./.pyrit/.env', './.pyrit/.env.local']\n",
      "Loaded environment file: ./.pyrit/.env\n",
      "Loaded environment file: ./.pyrit/.env.local\n",
      "gcg-romanlutz\n"
     ]
    }
   ],
   "source": [
    "import os\n",
    "\n",
    "from pyrit.setup.initialization import _load_environment_files\n",
    "\n",
    "_load_environment_files(env_files=None)\n",
    "\n",
    "subscription_id = os.environ.get(\"AZURE_ML_SUBSCRIPTION_ID\")\n",
    "resource_group = os.environ.get(\"AZURE_ML_RESOURCE_GROUP\")\n",
    "workspace = os.environ.get(\"AZURE_ML_WORKSPACE_NAME\")\n",
    "print(workspace)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "5",
   "metadata": {},
   "source": [
    "The Azure ML SDK emits a fair amount of telemetry to stderr that looks\n",
    "alarming but is benign: every operation logs an `ActivityCompleted: ...\n",
    "HowEnded=Failure` line for any expected `UserError` (such as\n",
    "`create_or_update` finding the environment already at the latest version),\n",
    "and every preview / experimental class prints a one-line warning. Quiet\n",
    "all of it so the rest of the notebook output stays focused on what\n",
    "actually matters."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "6",
   "metadata": {},
   "outputs": [],
   "source": [
    "import logging\n",
    "import warnings\n",
    "\n",
    "logging.getLogger(\"azure.ai.ml\").setLevel(logging.ERROR)\n",
    "warnings.filterwarnings(\"ignore\", module=r\"azure\\.ai\\.ml.*\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "7",
   "metadata": {},
   "outputs": [
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "Class DeploymentTemplateOperations: This is an experimental class, and may change at any time. Please see https://aka.ms/azuremlexperimental for more information.\n"
     ]
    }
   ],
   "source": [
    "from azure.ai.ml import MLClient\n",
    "from azure.identity import AzureCliCredential\n",
    "\n",
    "ml_client = MLClient(AzureCliCredential(), subscription_id, resource_group, workspace)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "8",
   "metadata": {},
   "source": [
    "## Create AML Environment"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "9",
   "metadata": {},
   "source": [
    "To install the dependencies needed to run GCG, we create an AML environment from a\n",
    "[Dockerfile](https://github.com/microsoft/PyRIT/blob/main/pyrit/executor/promptgen/gcg/src/Dockerfile). The Dockerfile uses\n",
    "an NVIDIA CUDA base image with Python 3.11 and installs PyRIT with the `gcg` extra."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "10",
   "metadata": {},
   "outputs": [
    {
     "data": {
      "text/plain": [
       "Environment({'arm_type': 'environment_version', 'latest_version': None, 'image': None, 'intellectual_property': None, 'is_anonymous': False, 'auto_increment_version': False, 'auto_delete_setting': None, 'name': 'pyrit-gcg', 'description': 'PyRIT GCG environment: CUDA 12.1 + Python 3.11 + pip install -e .[gcg]', 'tags': {'Owner': 'unknown'}, 'properties': {'azureml.labels': 'latest'}, 'print_as_yaml': False, 'id': '/subscriptions/db1ba766-2ca3-42c6-a19a-0f0d43134a8c/resourceGroups/gcg-romanlutz/providers/Microsoft.MachineLearningServices/workspaces/gcg-romanlutz/environments/pyrit-gcg/versions/17', 'Resource__source_path': '', 'base_path': './git/copilot-worktrees/PyRIT/romanlutz-upgraded-barnacle/doc/code/executor/gcg', 'creation_context': <azure.ai.ml.entities._system_data.SystemData object at 0x000002BBB8E0F770>, 'serialize': <msrest.serialization.Serializer object at 0x000002BBB9129320>, 'version': '17', 'conda_file': None, 'build': <azure.ai.ml.entities._assets.environment.BuildContext object at 0x000002BBB915FB10>, 'inference_config': None, 'os_type': 'Linux', 'conda_file_path': None, 'path': None, 'datastore': None, 'upload_hash': None, 'translated_conda_file': None})"
      ]
     },
     "execution_count": null,
     "metadata": {},
     "output_type": "execute_result"
    }
   ],
   "source": [
    "from pathlib import Path\n",
    "\n",
    "from azure.ai.ml.entities import BuildContext, Environment\n",
    "\n",
    "from pyrit.common.path import HOME_PATH\n",
    "\n",
    "# Configure the AML environment — build context is the repo root so the Dockerfile\n",
    "# can COPY pyproject.toml and pyrit/ for pip install -e \".[gcg]\"\n",
    "env_docker_context = Environment(\n",
    "    build=BuildContext(\n",
    "        path=Path(HOME_PATH),\n",
    "        dockerfile_path=\"pyrit/executor/promptgen/gcg/src/Dockerfile\",\n",
    "    ),\n",
    "    name=\"pyrit-gcg\",\n",
    "    description=\"PyRIT GCG environment: CUDA 12.1 + Python 3.11 + pip install -e .[gcg]\",\n",
    "    tags={\"Owner\": os.environ.get(\"USER\", \"unknown\")},\n",
    ")\n",
    "\n",
    "ml_client.environments.create_or_update(env_docker_context)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "11",
   "metadata": {},
   "source": [
    "## Submit Training Job to AML"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "12",
   "metadata": {},
   "source": [
    "Finally, we configure the command to run the GCG algorithm. The entry point is\n",
    "[`pyrit.executor.promptgen.gcg.experiments.run`](https://github.com/microsoft/PyRIT/blob/main/pyrit/executor/promptgen/gcg/experiments/run.py),\n",
    "invoked as a module so the uploaded code snapshot takes priority over the\n",
    "Docker-installed package (Python's `-m` flag puts the cwd at the front of `sys.path`).\n",
    "\n",
    "The new public API takes a typed ``GCGConfig`` (strategy) and a separate\n",
    "``GCGDataConfig`` (CSV paths/counts). We build both locally with whatever\n",
    "overrides we want, serialize each into a JSON file the AML job can read as\n",
    "an input, and ship those paths through the job command. Defaults come from\n",
    "the dataclasses in ``pyrit.executor.promptgen.gcg.config``; goals and targets\n",
    "flow into ``GCGGenerator.execute_async`` at runtime, not through the config.\n",
    "\n",
    "We also have to specify a GPU compute target. In our experience, a GPU instance with\n",
    "at least 24GB of vRAM is required (e.g., Standard_NC24ads_A100_v4).\n",
    "\n",
    "Depending on the compute instance you use, you may encounter \"out of memory\" errors.\n",
    "In this case, we recommend training on a smaller model or lowering ``data.n_train_data``\n",
    "or ``algorithm.batch_size``."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "13",
   "metadata": {},
   "outputs": [],
   "source": [
    "import tempfile\n",
    "\n",
    "from pyrit.executor.promptgen.gcg import (\n",
    "    GCGAlgorithmConfig,\n",
    "    GCGConfig,\n",
    "    GCGDataConfig,\n",
    "    GCGModelConfig,\n",
    "    GCGOutputConfig,\n",
    ")\n",
    "\n",
    "config = GCGConfig(\n",
    "    models=[GCGModelConfig(name=\"meta-llama/Llama-2-7b-chat-hf\")],\n",
    "    algorithm=GCGAlgorithmConfig(n_steps=5, batch_size=64, test_steps=1),\n",
    "    output=GCGOutputConfig(result_prefix=\"gcg_suffix\"),\n",
    ")\n",
    "data_config = GCGDataConfig(\n",
    "    train_data=(\"https://raw.githubusercontent.com/llm-attacks/llm-attacks/main/data/advbench/harmful_behaviors.csv\"),\n",
    "    n_train_data=5,\n",
    "    n_test_data=0,\n",
    ")\n",
    "\n",
    "# Write the configs into a tempdir so AML can mount them as separate job inputs.\n",
    "config_dir = Path(tempfile.mkdtemp(prefix=\"gcg-aml-config-\"))\n",
    "config_path = config_dir / \"config.json\"\n",
    "data_path = config_dir / \"data.json\"\n",
    "config.to_json_file(config_path)\n",
    "data_config.to_json_file(data_path)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "14",
   "metadata": {},
   "outputs": [],
   "source": [
    "from azure.ai.ml import Input, Output, command\n",
    "\n",
    "job = command(\n",
    "    code=Path(HOME_PATH),\n",
    "    command=(\n",
    "        \"python -m pyrit.executor.promptgen.gcg.experiments.run\"\n",
    "        \" --config ${{inputs.config}}\"\n",
    "        \" --data ${{inputs.data}}\"\n",
    "        \" --output-dir ${{outputs.results}}\"\n",
    "    ),\n",
    "    inputs={\n",
    "        \"config\": Input(type=\"uri_file\", path=str(config_path)),\n",
    "        \"data\": Input(type=\"uri_file\", path=str(data_path)),\n",
    "    },\n",
    "    outputs={\"results\": Output(type=\"uri_folder\")},\n",
    "    environment=f\"{env_docker_context.name}:{env_docker_context.version}\",\n",
    "    environment_variables={\"HUGGINGFACE_TOKEN\": os.environ[\"HUGGINGFACE_TOKEN\"]},\n",
    "    compute=\"gcg-gpu-a100\",\n",
    "    display_name=\"gcg_suffix_generation\",\n",
    "    description=\"Generate adversarial suffixes using GCG on Llama-2.\",\n",
    "    tags={\"Owner\": os.environ.get(\"USER\", \"unknown\")},\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "15",
   "metadata": {},
   "outputs": [
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "Class AutoDeleteSettingSchema: This is an experimental class, and may change at any time. Please see https://aka.ms/azuremlexperimental for more information.\n"
     ]
    },
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "Class AutoDeleteConditionSchema: This is an experimental class, and may change at any time. Please see https://aka.ms/azuremlexperimental for more information.\n"
     ]
    },
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "Class BaseAutoDeleteSettingSchema: This is an experimental class, and may change at any time. Please see https://aka.ms/azuremlexperimental for more information.\n"
     ]
    },
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "Class IntellectualPropertySchema: This is an experimental class, and may change at any time. Please see https://aka.ms/azuremlexperimental for more information.\n"
     ]
    },
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "Class ProtectionLevelSchema: This is an experimental class, and may change at any time. Please see https://aka.ms/azuremlexperimental for more information.\n"
     ]
    },
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "Class BaseIntellectualPropertySchema: This is an experimental class, and may change at any time. Please see https://aka.ms/azuremlexperimental for more information.\n"
     ]
    },
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "pathOnCompute is not a known attribute of class <class 'azure.ai.ml._restclient.v2023_04_01_preview.models._models_py3.UriFolderJobOutput'> and will be ignored\n"
     ]
    },
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "Job: great_vulture_lwn9y2fs10\n",
      "Status: Starting\n",
      "Studio URL: https://ml.azure.com/runs/great_vulture_lwn9y2fs10?wsid=/subscriptions/db1ba766-2ca3-42c6-a19a-0f0d43134a8c/resourcegroups/gcg-romanlutz/workspaces/gcg-romanlutz&tid=72f988bf-86f1-41af-91ab-2d7cd011db47\n"
     ]
    }
   ],
   "source": [
    "returned_job = ml_client.create_or_update(job)\n",
    "print(f\"Job: {returned_job.name}\")\n",
    "print(f\"Status: {returned_job.status}\")\n",
    "print(f\"Studio URL: {returned_job.studio_url}\")"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "16",
   "metadata": {},
   "source": [
    "## Wait for the Job to Complete and Inspect the Generated Suffix\n",
    "\n",
    "The next cell polls the job until it reaches a terminal state (~20-30\n",
    "minutes for the small 5-step baseline above), then downloads the named\n",
    "`results` output and prints the final suffix. The runner writes its\n",
    "result file as `<result_prefix>_<timestamp>.json` (with `result_prefix`\n",
    "coming from the `GCGConfig` we built above, plus the AML output mount\n",
    "prepended by `--output-dir`). For our config, that resolves to\n",
    "`gcg_suffix_<timestamp>.json` under\n",
    "`<download_dir>/named-outputs/results/` once we download. The\n",
    "`controls` array in that file contains one entry per training step, and\n",
    "the last entry is the final adversarial suffix that, appended to the user\n",
    "prompt, was optimized to elicit the target response."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "17",
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "Job status: Queued\n"
     ]
    },
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "Job status: Running\n"
     ]
    },
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "Job status: Completed\n"
     ]
    },
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "Downloading artifact azureml://subscriptions/db1ba766-2ca3-42c6-a19a-0f0d43134a8c/resourcegroups/gcg-romanlutz/workspaces/gcg-romanlutz/datastores/workspaceblobstore/paths/azureml/great_vulture_lwn9y2fs10/results/ to ./AppData/Local/Temp/gcg-aml-7ccxlx7y/named-outputs/results\n"
     ]
    },
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "Downloading artifact azureml://datastores/workspaceartifactstore/ExperimentRun/dcid.great_vulture_lwn9y2fs10 to ./AppData/Local/Temp/gcg-aml-7ccxlx7y/artifacts\n"
     ]
    },
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "Result file: gcg_suffix_20260523-051515.json\n",
      "Steps run: 30\n",
      "Final loss: 1.189453125\n",
      "Generated suffix: '! ! ! ! ! ! language ! programs ! ! Wisconsin ! ! ! ! Bill ! ! gender'\n"
     ]
    }
   ],
   "source": [
    "import json\n",
    "import tempfile\n",
    "import time\n",
    "from pathlib import Path\n",
    "\n",
    "_TERMINAL_STATES = {\"Completed\", \"Failed\", \"Canceled\", \"CancelRequested\"}\n",
    "\n",
    "last_status = None\n",
    "while True:\n",
    "    current_status = ml_client.jobs.get(returned_job.name).status\n",
    "    if current_status != last_status:\n",
    "        print(f\"Job status: {current_status}\", flush=True)\n",
    "        last_status = current_status\n",
    "    if current_status in _TERMINAL_STATES:\n",
    "        break\n",
    "    time.sleep(60)\n",
    "\n",
    "assert current_status == \"Completed\", f\"Job did not complete successfully: {current_status}\"\n",
    "\n",
    "download_dir = Path(tempfile.mkdtemp(prefix=\"gcg-aml-\"))\n",
    "ml_client.jobs.download(name=returned_job.name, download_path=str(download_dir), all=True)\n",
    "\n",
    "result_files = list(download_dir.rglob(\"gcg_suffix_*.json\"))\n",
    "if not result_files:\n",
    "    print(f\"No GCG result file found under {download_dir}. Files captured:\")\n",
    "    for p in sorted(download_dir.rglob(\"*\")):\n",
    "        if p.is_file():\n",
    "            print(f\"  {p.relative_to(download_dir)}\")\n",
    "    raise FileNotFoundError(\"Result JSON not in downloaded artifacts\")\n",
    "\n",
    "result_file = result_files[0]\n",
    "with open(result_file) as f:\n",
    "    log = json.load(f)\n",
    "\n",
    "final_suffix = log[\"controls\"][-1] if log[\"controls\"] else None\n",
    "final_loss = log[\"losses\"][-1] if log[\"losses\"] else None\n",
    "\n",
    "print(f\"Result file: {result_file.name}\")\n",
    "print(f\"Steps run: {len(log['controls'])}\")\n",
    "print(f\"Final loss: {final_loss}\")\n",
    "print(f\"Generated suffix: {final_suffix!r}\")"
   ]
  }
 ],
 "metadata": {
  "jupytext": {
   "cell_metadata_filter": "-all",
   "main_language": "python"
  },
  "language_info": {
   "codemirror_mode": {
    "name": "ipython",
    "version": 3
   },
   "file_extension": ".py",
   "mimetype": "text/x-python",
   "name": "python",
   "nbconvert_exporter": "python",
   "pygments_lexer": "ipython3",
   "version": "3.14.4"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
