{
  "cells": [
    {
      "cell_type": "code",
      "execution_count": null,
      "id": "7b93beeb",
      "metadata": {
        "id": "e706cd7e"
      },
      "outputs": [],
      "source": [
        "# Copyright 2025 Google LLC\n",
        "#\n",
        "# Licensed under the Apache License, Version 2.0 (the \"License\");\n",
        "# you may not use this file except in compliance with the License.\n",
        "# You may obtain a copy of the License at\n",
        "#\n",
        "#     https://www.apache.org/licenses/LICENSE-2.0\n",
        "#\n",
        "# Unless required by applicable law or agreed to in writing, software\n",
        "# distributed under the License is distributed on an \"AS IS\" BASIS,\n",
        "# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n",
        "# See the License for the specific language governing permissions and\n",
        "# limitations under the License."
      ]
    },
    {
      "cell_type": "markdown",
      "id": "0ef3b9b5",
      "metadata": {
        "id": "6302e5e3"
      },
      "source": [
        "# Intro to Agent Platform Multimodal Datasets\n",
        "\n",
        "<table align=\"left\">\n",
        "  <td style=\"text-align: center\">\n",
        "    <a href=\"https://colab.research.google.com/github/GoogleCloudPlatform/generative-ai/blob/main/gemini/multimodal-dataset/intro_vertex_ai_multimodal_dataset.ipynb\">\n",
        "      <img width=\"32px\" src=\"https://www.gstatic.com/pantheon/images/bigquery/welcome_page/colab-logo.svg\" alt=\"Google Colaboratory logo\"><br> Open in Colab\n",
        "    </a>\n",
        "  </td>\n",
        "  <td style=\"text-align: center\">\n",
        "    <a href=\"https://console.cloud.google.com/vertex-ai/colab/import/https:%2F%2Fraw.githubusercontent.com%2FGoogleCloudPlatform%2Fgenerative-ai%2Fmain%2Fgemini%2Fmultimodal-dataset%2Fintro_vertex_ai_multimodal_dataset.ipynb\">\n",
        "      <img width=\"32px\" src=\"https://lh3.googleusercontent.com/JmcxdQi-qOpctIvWKgPtrzZdJJK-J3sWE1RsfjZNwshCFgE_9fULcNpuXYTilIR2hjwN\" alt=\"Google Cloud Colab Enterprise logo\"><br> Open in Colab Enterprise\n",
        "    </a>\n",
        "  </td>\n",
        "  <td style=\"text-align: center\">\n",
        "    <a href=\"https://console.cloud.google.com/vertex-ai/workbench/deploy-notebook?download_url=https://raw.githubusercontent.com/GoogleCloudPlatform/generative-ai/main/gemini/multimodal-dataset/intro_vertex_ai_multimodal_dataset.ipynb\">\n",
        "      <img src=\"https://www.gstatic.com/images/branding/gcpiconscolors/vertexai/v1/32px.svg\" alt=\"Agent Platform logo\"><br> Open in Agent Platform Workbench\n",
        "    </a>\n",
        "  </td>\n",
        "  <td style=\"text-align: center\">\n",
        "    <a href=\"https://console.cloud.google.com/bigquery/import?url=https://github.com/GoogleCloudPlatform/generative-ai/blob/main/gemini/multimodal-dataset/intro_vertex_ai_multimodal_dataset.ipynb\">\n",
        "      <img src=\"https://www.gstatic.com/images/branding/gcpiconscolors/bigquery/v1/32px.svg\" alt=\"BigQuery Studio logo\"><br> Open in BigQuery Studio\n",
        "    </a>\n",
        "  </td>\n",
        "  <td style=\"text-align: center\">\n",
        "    <a href=\"https://github.com/GoogleCloudPlatform/generative-ai/blob/main/gemini/multimodal-dataset/intro_vertex_ai_multimodal_dataset.ipynb\">\n",
        "      <img width=\"32px\" src=\"https://raw.githubusercontent.com/primer/octicons/refs/heads/main/icons/mark-github-24.svg\" alt=\"GitHub logo\"><br> View on GitHub\n",
        "    </a>\n",
        "  </td>\n",
        "</table>\n",
        "\n",
        "<div style=\"clear: both;\"></div>\n",
        "\n",
        "<b>Share to:</b>\n",
        "\n",
        "<a href=\"https://www.linkedin.com/sharing/share-offsite/?url=https%3A//github.com/GoogleCloudPlatform/generative-ai/blob/main/gemini/multimodal-dataset/intro_vertex_ai_multimodal_dataset.ipynb\" target=\"_blank\">\n",
        "  <img width=\"20px\" src=\"https://upload.wikimedia.org/wikipedia/commons/8/81/LinkedIn_icon.svg\" alt=\"LinkedIn logo\">\n",
        "</a>\n",
        "\n",
        "<a href=\"https://bsky.app/intent/compose?text=https%3A//github.com/GoogleCloudPlatform/generative-ai/blob/main/gemini/multimodal-dataset/intro_vertex_ai_multimodal_dataset.ipynb\" target=\"_blank\">\n",
        "  <img width=\"20px\" src=\"https://upload.wikimedia.org/wikipedia/commons/7/7a/Bluesky_Logo.svg\" alt=\"Bluesky logo\">\n",
        "</a>\n",
        "\n",
        "<a href=\"https://twitter.com/intent/tweet?url=https%3A//github.com/GoogleCloudPlatform/generative-ai/blob/main/gemini/multimodal-dataset/intro_vertex_ai_multimodal_dataset.ipynb\" target=\"_blank\">\n",
        "  <img width=\"20px\" src=\"https://upload.wikimedia.org/wikipedia/commons/5/5a/X_icon_2.svg\" alt=\"X logo\">\n",
        "</a>\n",
        "\n",
        "<a href=\"https://reddit.com/submit?url=https%3A//github.com/GoogleCloudPlatform/generative-ai/blob/main/gemini/multimodal-dataset/intro_vertex_ai_multimodal_dataset.ipynb\" target=\"_blank\">\n",
        "  <img width=\"20px\" src=\"https://redditinc.com/hubfs/Reddit%20Inc/Brand/Reddit_Logo.png\" alt=\"Reddit logo\">\n",
        "</a>\n",
        "\n",
        "<a href=\"https://www.facebook.com/sharer/sharer.php?u=https%3A//github.com/GoogleCloudPlatform/generative-ai/blob/main/gemini/multimodal-dataset/intro_vertex_ai_multimodal_dataset.ipynb\" target=\"_blank\">\n",
        "  <img width=\"20px\" src=\"https://upload.wikimedia.org/wikipedia/commons/5/51/Facebook_f_logo_%282019%29.svg\" alt=\"Facebook logo\">\n",
        "</a>"
      ]
    },
    {
      "cell_type": "markdown",
      "id": "c03e2730",
      "metadata": {
        "id": "8ac41575"
      },
      "source": [
        "| Authors |\n",
        "| --- |\n",
        "| [Frances Thoma](https://github.com/diskontinuum) |\n",
        "| [Christian Leopoldseder](https://github.com/cleop-google) |"
      ]
    },
    {
      "cell_type": "markdown",
      "id": "26760bb6",
      "metadata": {
        "id": "327042bb"
      },
      "source": [
        "## Overview\n",
        "\n",
        "This notebook demonstrates how to use Agent Platform Multimodal Datasets to assemble Gemini requests, to run a validation and resource estimation for supervised fine-tuning, and to create tuning and batch prediction jobs.\n",
        "\n",
        "### Objectives\n",
        "\n",
        "- Preview the new Agent Platform Multimodal Datasets SDK\n",
        "- Demo upcoming integrations\n",
        "\n",
        "### Costs\n",
        "\n",
        "This tutorial uses billable components of Google Cloud:\n",
        "\n",
        "* Agent Platform\n",
        "* Cloud Storage\n",
        "* BigQuery\n",
        "\n",
        "Learn about [Agent Platform pricing](https://cloud.google.com/vertex-ai/pricing), [Cloud Storage pricing](https://cloud.google.com/storage/pricing), and [BigQuery pricing](https://cloud.google.com/bigquery/pricing) and use the [Pricing Calculator](https://cloud.google.com/products/calculator/) to generate a cost estimate based on your projected usage.\n",
        "\n",
        "### Prerequisites\n",
        "1. Make sure that [billing is enabled](https://cloud.google.com/billing/docs/how-to/modify-project) for your project.\n",
        "\n",
        "2. You must have an existing Google Cloud project and [enable the Agent Platform API](https://console.cloud.google.com/flows/enableapi?apiid=aiplatform.googleapis.com). Learn more about [setting up a project and a development environment](https://cloud.google.com/vertex-ai/docs/start/cloud-environment).\n",
        "\n",
        "### Questions or Feedback\n",
        "\n",
        "You can reach out directly to the authors via `vertex-multimodal-dataset-external-feedback@google.com` for feedback or questions."
      ]
    },
    {
      "cell_type": "markdown",
      "id": "c4faee94",
      "metadata": {
        "id": "a6a32c38"
      },
      "source": [
        "## Get Started"
      ]
    },
    {
      "cell_type": "markdown",
      "id": "8a387418",
      "metadata": {
        "id": "fd0ebd94"
      },
      "source": [
        "### Install Agent Platform SDK and other required packages"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "id": "ba48aa45",
      "metadata": {
        "id": "b28f0ace"
      },
      "outputs": [],
      "source": [
        "%pip install --quiet --upgrade google-cloud-aiplatform google-genai bigframes"
      ]
    },
    {
      "cell_type": "markdown",
      "id": "b5c16e04",
      "metadata": {
        "id": "8e69fb2b"
      },
      "source": [
        "### Authenticate your notebook environment (Colab only)\n",
        "\n",
        "If you are running this notebook on Google Colab, run the cell below to authenticate your environment."
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "id": "ab34ead8",
      "metadata": {
        "id": "25681e6e"
      },
      "outputs": [],
      "source": [
        "import sys\n",
        "\n",
        "if \"google.colab\" in sys.modules:\n",
        "    from google.colab import auth\n",
        "\n",
        "    auth.authenticate_user()"
      ]
    },
    {
      "cell_type": "markdown",
      "id": "a470635a",
      "metadata": {
        "id": "672cc818"
      },
      "source": [
        "- If you are running this notebook in a local development environment:\n",
        "  - Install the [Google Cloud SDK](https://cloud.google.com/sdk).\n",
        "  - Obtain authentication credentials. Create local credentials by running the following command and following the oauth2 flow (read more about the command [here](https://cloud.google.com/sdk/gcloud/reference/beta/auth/application-default/login)):\n",
        "\n",
        "    ```bash\n",
        "    gcloud auth application-default login\n",
        "    ```"
      ]
    },
    {
      "cell_type": "markdown",
      "id": "e6287833",
      "metadata": {
        "id": "c6bb8d65"
      },
      "source": [
        "### Import libraries"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "id": "3120f466",
      "metadata": {
        "id": "98fb079d"
      },
      "outputs": [],
      "source": [
        "import io\n",
        "import os\n",
        "\n",
        "import agentplatform\n",
        "import bigframes.pandas as bpd\n",
        "import pandas\n",
        "from PIL import Image\n",
        "from agentplatform.types import (\n",
        "    GeminiExample,\n",
        "    GeminiRequestReadConfig,\n",
        "    GeminiTemplateConfig,\n",
        ")\n",
        "from google import genai\n",
        "from google.cloud import storage\n",
        "from google.genai.types import Content, CreateTuningJobConfig, JobState, Part"
      ]
    },
    {
      "cell_type": "markdown",
      "id": "c640a9eb",
      "metadata": {
        "id": "0b8bc486"
      },
      "source": [
        "### Set Google Cloud project information\n",
        "\n",
        "To get started using Agent Platform, you must have an existing Google Cloud project and [enable the Agent Platform API](https://console.cloud.google.com/flows/enableapi?apiid=aiplatform.googleapis.com).\n",
        "\n",
        "Learn more about [setting up a project and a development environment](https://cloud.google.com/vertex-ai/docs/start/cloud-environment)."
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "id": "d29d9ef1",
      "metadata": {
        "id": "35c835bc"
      },
      "outputs": [],
      "source": [
        "# Use the environment variable if the user doesn't provide Project ID.\n",
        "# fmt: off\n",
        "PROJECT_ID = \"[your-project-id]\"  # @param {type: \"string\", placeholder: \"[your-project-id]\", isTemplate: true}\n",
        "# fmt: on\n",
        "if not PROJECT_ID or PROJECT_ID == \"[your-project-id]\":\n",
        "    PROJECT_ID = str(os.environ.get(\"GOOGLE_CLOUD_PROJECT\"))\n",
        "\n",
        "LOCATION = os.environ.get(\"GOOGLE_CLOUD_REGION\", \"us-central1\")\n",
        "\n",
        "client = agentplatform.Client(project=PROJECT_ID, location=LOCATION)\n",
        "\n",
        "# BigFrames settings\n",
        "bpd.close_session()\n",
        "bpd.options.bigquery.project = PROJECT_ID\n",
        "bpd.options.bigquery.location = LOCATION"
      ]
    },
    {
      "cell_type": "markdown",
      "id": "7c48878c",
      "metadata": {
        "id": "c2130a7f"
      },
      "source": [
        "### Data preparation\n",
        "\n",
        "The image files and labels used in this tutorial are from the flower dataset used in this [TensorFlow blog post](https://cloud.google.com/blog/products/gcp/how-to-classify-images-with-tensorflow-using-google-cloud-machine-learning-and-cloud-dataflow).\n",
        "\n",
        "The dataset contains 7338 images, each of which is annotated with one label across 5 different flower classes.\n",
        "\n",
        "The input images are stored in a public Cloud Storage bucket. This publicly-accessible bucket also contains a CSV file used to create the Agent Platform multimodal dataset. This file has two columns: the first column lists an image's URI in Cloud Storage, and the second column contains the image's label.\n",
        "\n",
        "In this notebook, we'll use subsets of the flower dataset, each with a fixed number of examples per category, and prepare training, tuning and test subsets\n",
        " as DataFrame.\n",
        "\n",
        "**Tip:** Use the BigFrames library `bpd` instead of `pandas` for larger datasets."
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "id": "c73b3ae1",
      "metadata": {
        "id": "3f7a663b"
      },
      "outputs": [],
      "source": [
        "# Get data from GCS\n",
        "csv = \"gs://cloud-samples-data/ai-platform/flowers/flowers.csv\"\n",
        "all_images = pandas.read_csv(csv, names=[\"image_uris\", \"labels\"])\n",
        "\n",
        "# Shuffle\n",
        "all_images = all_images.sample(frac=1).reset_index(drop=True)\n",
        "\n",
        "# Prepare training and validation set\n",
        "CATEGORIES = [\"daisy\", \"dandelion\", \"roses\", \"sunflowers\", \"tulips\"]\n",
        "# fmt: off\n",
        "TRAINING_CASES_PER_CATEGORY = 100  # @param {type: 'integer'}\n",
        "VALIDATION_CASES_PER_CATEGORY = 100  # @param {type: 'integer'}\n",
        "PREDICTION_DATASET_SIZE = 100  # @param {type: 'integer'}\n",
        "# fmt: on\n",
        "\n",
        "# Set up the prediction set\n",
        "if len(all_images) < PREDICTION_DATASET_SIZE:\n",
        "    raise ValueError(\n",
        "        \"Prediction dataset size is larger than the total number of images.\"\n",
        "    )\n",
        "prediction_set = all_images.iloc[:PREDICTION_DATASET_SIZE]\n",
        "all_images = all_images.iloc[PREDICTION_DATASET_SIZE:]\n",
        "\n",
        "# Set up the training and validation set with evenly distributed labels\n",
        "training_set = pandas.DataFrame()\n",
        "validation_set = pandas.DataFrame()\n",
        "\n",
        "\n",
        "for category in CATEGORIES:\n",
        "    same_labels = all_images[all_images[\"labels\"] == category]\n",
        "    if len(same_labels) < TRAINING_CASES_PER_CATEGORY + VALIDATION_CASES_PER_CATEGORY:\n",
        "        raise ValueError(\"Please reduce the number of cases per category.\")\n",
        "    training_set = pandas.concat(\n",
        "        (training_set, same_labels.iloc[:TRAINING_CASES_PER_CATEGORY]),\n",
        "        ignore_index=True,\n",
        "    )\n",
        "    validation_set = pandas.concat(\n",
        "        (\n",
        "            validation_set,\n",
        "            same_labels.iloc[\n",
        "                TRAINING_CASES_PER_CATEGORY : TRAINING_CASES_PER_CATEGORY\n",
        "                + VALIDATION_CASES_PER_CATEGORY\n",
        "            ],\n",
        "        ),\n",
        "        ignore_index=True,\n",
        "    )"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "id": "72b3d7d7",
      "metadata": {
        "cellView": "form",
        "id": "7dd7943d"
      },
      "outputs": [],
      "source": [
        "# @title Common Functions\n",
        "\n",
        "# Set Pandas display options to show all columns and full width for better inspection\n",
        "pandas.set_option(\"display.max_columns\", None)  # Show all columns\n",
        "pandas.set_option(\"display.expand_frame_repr\", False)  # Prevent line wrapping\n",
        "pandas.set_option(\"display.max_colwidth\", None)  # Show full column width\n",
        "\n",
        "\n",
        "def show_dataset_info(dataset):\n",
        "    \"\"\"Dataset inspection helper\"\"\"\n",
        "    print(\"  Resource name: \", dataset.name)\n",
        "    print(\"  Display name: \", dataset.display_name)\n",
        "    print(\"  BQ Table:     \", dataset.bigquery_uri)\n",
        "\n",
        "\n",
        "def get_gcs_image(gcs_uri):\n",
        "    \"\"\"Download and show an image from Cloud Storage.\"\"\"\n",
        "    storage_client = storage.Client(project=PROJECT_ID)\n",
        "    blob = storage.blob.Blob.from_string(gcs_uri, client=storage_client)\n",
        "    return Image.open(io.BytesIO(blob.download_as_bytes()))"
      ]
    },
    {
      "cell_type": "markdown",
      "id": "3fafaae1",
      "metadata": {
        "id": "103f2d4c"
      },
      "source": [
        "## User Journey Demo\n",
        "\n",
        "The user journey demonstrated here contains the following steps:\n",
        "\n",
        "1. Create Dataset\n",
        "2. Assemble the dataset with a template and inspect assembly\n",
        "3. Run a validation for tuning\n",
        "4. Estimate Resources for tuning\n",
        "5. Run tuning\n",
        "6. Run batch prediction"
      ]
    },
    {
      "cell_type": "markdown",
      "id": "73f0c5ed",
      "metadata": {
        "id": "c4617830"
      },
      "source": [
        "### 1. Create a dataset from a Pandas or BigFrames DataFrame"
      ]
    },
    {
      "cell_type": "markdown",
      "id": "a6f543cd",
      "metadata": {
        "id": "96c6368d"
      },
      "source": [
        "We prepared a DataFrame `training_set` with two columns:\n",
        "\n",
        "*   `image_uris`: GCS URIs of flower images\n",
        "*   `labels`: Flower label (five flower categories, one label per image)"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "id": "85492029",
      "metadata": {
        "id": "c4152b17"
      },
      "outputs": [],
      "source": [
        "flower_uri = training_set[\"image_uris\"].iloc[0]\n",
        "flower_label = training_set[\"labels\"].iloc[0]\n",
        "\n",
        "display(get_gcs_image(flower_uri))\n",
        "print(f\"Image URI: {flower_uri}\")\n",
        "print(f\"Flower label: {flower_label}\")\n",
        "training_set.head()"
      ]
    },
    {
      "cell_type": "markdown",
      "id": "063e4fcf",
      "metadata": {
        "id": "810e525f"
      },
      "source": [
        "Let's create a multimodal dataset from the prepared DataFrame."
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "id": "dc6e1a84",
      "metadata": {
        "id": "a2467da5"
      },
      "outputs": [],
      "source": [
        "flowers = client.datasets.create_from_pandas(dataframe=training_set)\n",
        "\n",
        "# Inspect the dataset\n",
        "show_dataset_info(flowers)\n",
        "flowers.to_bigframes().head()"
      ]
    },
    {
      "cell_type": "markdown",
      "id": "94af2e34",
      "metadata": {
        "id": "66f39229"
      },
      "source": [
        "**Other dataset creation options**\n",
        "\n",
        "Create from a BigQuery table (pass the URI directly).\n",
        "\n",
        "```py\n",
        "my_dataset_from_bigquery = client.datasets.create_from_bigquery(\n",
        "    bigquery_uri=\"bq://projectId.datasetId.tableId\",\n",
        ")\n",
        "```\n",
        "\n",
        "Create from a BigFrames DataFrame.\n",
        "\n",
        "```py\n",
        "my_dataset_from_bigframes = client.datasets.create_from_bigframes(\n",
        "    dataframe=my_dataframe,\n",
        ")\n",
        "```\n",
        "\n",
        "Create from a GCS file in JSONL format, where each line is an assembled Gemini `GenerateContentRequest` (no assembly step needed). The requests are stored in a single `requests` column in the backing BigQuery table.\n",
        "\n",
        "```py\n",
        "my_dataset_from_jsonl = client.datasets.create_from_gemini_request_jsonl(\n",
        "    gcs_uri=\"gs://my-bucket/path/to/requests.jsonl\",\n",
        ")\n",
        "```\n",
        "\n",
        "Load an existing dataset.\n",
        "\n",
        "```py\n",
        "# Load dataset based on dataset name (accepts full resource name or dataset ID)\n",
        "same_dataset = client.datasets.get_multimodal_dataset(name=dataset_name)\n",
        "```"
      ]
    },
    {
      "cell_type": "markdown",
      "id": "409f309a",
      "metadata": {
        "id": "644b0f2e"
      },
      "source": [
        "### 2. Assemble the dataset with a template and inspect assembly\n",
        "\n",
        "To use our Flowers dataset with Gemini, let's assemble a full Gemini request referencing the images in our dataset."
      ]
    },
    {
      "cell_type": "markdown",
      "id": "bd9e1861",
      "metadata": {
        "id": "d3dd451c"
      },
      "source": [
        "We construct a read configuration using `GeminiRequestReadConfig.single_turn_template()` by specifying the general prompt, response and system instructions and use placeholders in curly braces. During the assembly, the placeholders are replaced with the values of the dataset column that the placeholders denote.\n",
        "\n",
        "The dataset columns referenced by the placeholders can contain e.g. GCS URIS for files of several data types and modalities:\n",
        "- .pdf\n",
        "- .png, .jpeg, .jpg, .webp\n",
        "- .aac, .flac, .mp3, .m4a, .mpga, .opus, .pcm, .wav\n",
        "- .flv, .mov, .mpegps, .mpg, .wmv, .3pg"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "id": "faa9c24f",
      "metadata": {
        "id": "b6a90bf6"
      },
      "outputs": [],
      "source": [
        "read_config = GeminiRequestReadConfig.single_turn_template(\n",
        "    prompt=\"This is the image: {image_uris}\",\n",
        "    response=\"{labels}\",\n",
        "    system_instruction=\"You are a botanical image classifier. Analyze the provided image \"\n",
        "    \"and determine the most accurate classification of the flower.\"\n",
        "    'These are the only flower categories: [/\"daisy/\", /\"dandelion/\", /\"roses/\", /\"sunflowers/\", /\"tulips/\"].'\n",
        "    \"Return only one category per image.\",\n",
        ")"
      ]
    },
    {
      "cell_type": "markdown",
      "id": "3e962188",
      "metadata": {
        "id": "077af951"
      },
      "source": [
        "Here, the read configuration is constructed using the `GeminiRequestReadConfig.single_turn_template()` class method. Alternatively, it can be explicitly constructed from a Gemini example as below.\n",
        "\n",
        "It is also possible to specify a custom field mapping for the placeholders used in the Gemini example. Then the placeholders can have any name, and not necessarily the column name of the dataset column with the values that are being inserted (here image_uris and labels):\n",
        "\n",
        "```\n",
        "gemini_example = GeminiExample(\n",
        "    contents=[\n",
        "        Content(role=\"user\", parts=[Part.from_text(text=\"This is the image: {uri}\")]),\n",
        "        Content(role=\"model\", parts=[Part.from_text(text=\"{flower}\")]),\n",
        "    ],\n",
        "    system_instruction=Content(\n",
        "        parts=[\n",
        "            Part.from_text(\n",
        "                text='You are a botanical image classifier. Analyze the provided image '\n",
        "                'and determine the most accurate classification of the flower.'\n",
        "                'These are the only flower categories: [/\"daisy/\", /\"dandelion/\", /\"roses/\", /\"sunflowers/\", /\"tulips/\"].'\n",
        "                'Return only one category per image.'\n",
        "            )\n",
        "        ]\n",
        "    ),\n",
        ")\n",
        "\n",
        "read_config = GeminiRequestReadConfig(\n",
        "    template_config=GeminiTemplateConfig(\n",
        "        gemini_example=gemini_example,\n",
        "        field_mapping={\"uri\": \"image_uris\", \"flower\": \"labels\"},\n",
        "    ),\n",
        ")\n",
        "```"
      ]
    },
    {
      "cell_type": "markdown",
      "id": "2b901de9",
      "metadata": {
        "id": "fb1babd8"
      },
      "source": [
        "**Assemble and inspect the dataset.**\n",
        "\n",
        "The dataset assembly creates a BigQuery table with the assembled examples in a single `request` column. The assembly method returns a `(table_id, dataframe)` tuple; pass `load_dataframe=True` to also load the assembled table as a BigFrames DataFrame for inspection."
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "id": "eae7a392",
      "metadata": {
        "id": "eb085788"
      },
      "outputs": [],
      "source": [
        "table_id, assembly = client.datasets.assemble(\n",
        "    name=flowers.name,\n",
        "    gemini_request_read_config=read_config,\n",
        "    load_dataframe=True,\n",
        ")\n",
        "\n",
        "# Inspect assembled dataset\n",
        "assembly.head()"
      ]
    },
    {
      "cell_type": "markdown",
      "id": "b07cf363",
      "metadata": {
        "id": "2d679b00"
      },
      "source": [
        "It is also possible to attach the read configuration to the dataset and run the assembly and the validation below without passing it.\n",
        "\n",
        "The read config is required for tuning and batch prediction jobs, so let's attach it now."
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "id": "d95344b8",
      "metadata": {
        "id": "4c4de92d"
      },
      "outputs": [],
      "source": [
        "flowers.set_read_config(read_config=read_config)\n",
        "flowers = client.datasets.update_multimodal_dataset(multimodal_dataset=flowers)\n",
        "_ = client.datasets.assemble(name=flowers.name)"
      ]
    },
    {
      "cell_type": "markdown",
      "id": "149b9440",
      "metadata": {
        "id": "6d6937ed"
      },
      "source": [
        "### 3. Run a validation for tuning"
      ]
    },
    {
      "cell_type": "markdown",
      "id": "50e5bc45",
      "metadata": {
        "id": "ca4baeb8"
      },
      "source": [
        "Validate a dataset for tuning.\n",
        "Tuning dataset usages are: `SFT_VALIDATION`, `SFT_TRAINING`.\n",
        "\n",
        "Since we attached the `read_config` above, it is used implicitly."
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "id": "14c70f6f",
      "metadata": {
        "id": "15819b9d"
      },
      "outputs": [],
      "source": [
        "validation = client.datasets.assess_tuning_validity(\n",
        "    dataset_name=flowers.name,\n",
        "    model_name=\"gemini-2.0-flash-001\",\n",
        "    dataset_usage=\"SFT_TRAINING\",\n",
        ")\n",
        "\n",
        "# Check if there are validation errors\n",
        "validation.errors"
      ]
    },
    {
      "cell_type": "markdown",
      "id": "36dbd143",
      "metadata": {
        "id": "d35d678b"
      },
      "source": [
        "Let's validate a dataset with an incorrect read configuration, e.g. using a `GeminiExample` that contains two consecutive `user` contents, instead of a `user` content followed by a `model` content."
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "id": "a320f575",
      "metadata": {
        "id": "c5694f52"
      },
      "outputs": [],
      "source": [
        "invalid_gemini_example = GeminiExample(\n",
        "    contents=[\n",
        "        Content(\n",
        "            role=\"user\",\n",
        "            parts=[Part.from_text(text=\"This is the image: {image_uris}\")],\n",
        "        ),\n",
        "        # Consecutive content turn with the same role\n",
        "        Content(\n",
        "            role=\"user\",\n",
        "            parts=[Part.from_text(text=\".\")],\n",
        "        ),\n",
        "    ],\n",
        ")\n",
        "invalid_read_config = GeminiRequestReadConfig(\n",
        "    template_config=GeminiTemplateConfig(gemini_example=invalid_gemini_example)\n",
        ")\n",
        "\n",
        "validation = client.datasets.assess_tuning_validity(\n",
        "    dataset_name=flowers.name,\n",
        "    model_name=\"gemini-2.0-flash-001\",\n",
        "    dataset_usage=\"SFT_TRAINING\",\n",
        "    gemini_request_read_config=invalid_read_config,\n",
        ")\n",
        "\n",
        "validation.errors"
      ]
    },
    {
      "cell_type": "markdown",
      "id": "0d90fac8",
      "metadata": {
        "id": "b994da9a"
      },
      "source": [
        "### 4. Estimate resources for tuning"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "id": "713d66b2",
      "metadata": {
        "id": "28bd43f1"
      },
      "outputs": [],
      "source": [
        "tuning_resources = client.datasets.assess_tuning_resources(\n",
        "    dataset_name=flowers.name, model_name=\"gemini-2.5-flash-001\"\n",
        ")\n",
        "print(tuning_resources)"
      ]
    },
    {
      "cell_type": "markdown",
      "id": "3417569f",
      "metadata": {
        "id": "fcf3863c"
      },
      "source": [
        "## 5. Run Tuning\n",
        "\n",
        "Prerequisites:\n",
        "\n",
        "- Your Agent Platform service account `service-{project_number}@gcp-sa-vertex-tune.iam.gserviceaccount.com` needs to have read permissions on to the GCS buckets referenced in the dataset. If this is not automatically the case then you might have to assign the service account a role such as `Storage Object User`, see the screenshot below.\n",
        "\n",
        "- The multimodal dataset needs to have an attached read configuration (see `set_read_config` + `update_multimodal_dataset` above)."
      ]
    },
    {
      "cell_type": "markdown",
      "id": "cc312c6d",
      "metadata": {
        "id": "9e8b6cf5"
      },
      "source": [
        "Let's also prepare and use the validation dataset."
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "id": "b3bab9a6",
      "metadata": {
        "id": "a94a6aca"
      },
      "outputs": [],
      "source": [
        "# Create multimodal dataset for the validation set and attach read config\n",
        "flowers_validation = client.datasets.create_from_pandas(dataframe=validation_set)\n",
        "\n",
        "flowers_validation.set_read_config(read_config=read_config)\n",
        "flowers_validation = client.datasets.update_multimodal_dataset(\n",
        "    multimodal_dataset=flowers_validation\n",
        ")"
      ]
    },
    {
      "cell_type": "markdown",
      "id": "2c40ebbb",
      "metadata": {
        "id": "cfaa7d64"
      },
      "source": [
        "Here we use the training and validation set to start a tuning job:\n"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "id": "23e6f015",
      "metadata": {
        "id": "41c2d84e"
      },
      "outputs": [],
      "source": [
        "genai_client = genai.Client(vertexai=True, project=PROJECT_ID, location=LOCATION)\n",
        "\n",
        "tuning_job = genai_client.tunings.tune(\n",
        "    base_model=\"gemini-2.5-flash\",\n",
        "    training_dataset={\"vertex_dataset_resource\": flowers.name},\n",
        "    config=CreateTuningJobConfig(\n",
        "        validation_dataset={\"vertex_dataset_resource\": flowers_validation.name},\n",
        "    ),\n",
        ")"
      ]
    },
    {
      "cell_type": "markdown",
      "id": "95c81d32",
      "metadata": {
        "id": "a87c2d26"
      },
      "source": [
        "Let's monitor the job state and obtain the model ID of the tuned model once the tuning job has ended."
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "id": "2d5a93ae",
      "metadata": {
        "id": "dcd4ddfc"
      },
      "outputs": [],
      "source": [
        "import time\n",
        "\n",
        "print(f\"Tuning job started: {tuning_job.name}\")\n",
        "\n",
        "# Wait for the job to complete\n",
        "while tuning_job.state not in (\n",
        "    \"JOB_STATE_SUCCEEDED\",\n",
        "    \"JOB_STATE_FAILED\",\n",
        "    \"JOB_STATE_CANCELLED\",\n",
        "):\n",
        "    time.sleep(60)\n",
        "    tuning_job = genai_client.tunings.get(name=tuning_job.name)\n",
        "    print(f\"Polling - Current job state: {tuning_job.state}\")\n",
        "\n",
        "# Check the final state\n",
        "if tuning_job.state == \"JOB_STATE_FAILED\":\n",
        "    print(f\"Tuning job failed: {tuning_job.error}\")\n",
        "else:\n",
        "    print(f\"Tuning job ended with state: {tuning_job.state}\")\n",
        "    # Get model ID\n",
        "    tuned_model_id = tuning_job.tuned_model.model.split(\"/\")[-1]"
      ]
    },
    {
      "cell_type": "markdown",
      "id": "70c582f4",
      "metadata": {
        "id": "10ad0971"
      },
      "source": [
        "## 6. Batch Prediction\n",
        "\n",
        "You can use Multimodal datasets to run a batch prediction job with Gemini.\n",
        "You can pass the dataset resource name directly as the batch job source — no manual assembly step needed:\n",
        "\n",
        "```py\n",
        "job = genai_client.batches.create(\n",
        "    model=model,\n",
        "    src=flowers_prediction.name,\n",
        ")\n",
        "```\n",
        "\n",
        "Alternatively, you can still assemble the dataset manually and pass the assembled BigQuery table:\n",
        "\n",
        "```py\n",
        "table_id, _ = client.datasets.assemble(\n",
        "    name=flowers_prediction.name,\n",
        "    gemini_request_read_config=read_config,\n",
        ")\n",
        "\n",
        "job = genai_client.batches.create(\n",
        "    model=model,\n",
        "    src=f\"bq://{table_id}\",\n",
        ")\n",
        "```\n",
        "\n",
        "See also: [Batch Prediction Documentation](https://cloud.google.com/vertex-ai/generative-ai/docs/multimodal/batch-prediction-gemini).\n"
      ]
    },
    {
      "cell_type": "markdown",
      "id": "41a92f14",
      "metadata": {
        "id": "67baf2e8"
      },
      "source": [
        "The model specified can either be a Gemini base model, the model you tuned above, or any other tuned custom model.\n",
        "\n",
        "By default the colab uses the model tuned above for the batch prediction. If you want to run a batch prediction job on a base model or on another custom model, you can provide the base model name or the custom model id, respectively, in the field below."
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "id": "d12ffb2a",
      "metadata": {
        "cellView": "form",
        "id": "fd552d27"
      },
      "outputs": [],
      "source": [
        "# @title Specify the Model used for Batch Prediction\n",
        "# fmt: off\n",
        "model_for_batch_prediction = \"tuned model\"  # @param [\"tuned model\", \"base model\", \"other custom model\"]\n",
        "# fmt: on\n",
        "optional_base_model_name = \"\"  # @param {type:\"string\"}\n",
        "optional_custom_model_id = \"\"  # @param {type:\"string\"}\n",
        "\n",
        "# Set full model name\n",
        "if model_for_batch_prediction == \"tuned model\":\n",
        "    model = f\"projects/{PROJECT_ID}/locations/{LOCATION}/models/{tuned_model_id}\"\n",
        "elif model_for_batch_prediction == \"base model\":\n",
        "    if not optional_base_model_name:\n",
        "        raise ValueError(\n",
        "            \"Please provide a optional_base_model_name when 'base model' is selected.\"\n",
        "        )\n",
        "    model = optional_base_model_name\n",
        "elif model_for_batch_prediction == \"other custom model\":\n",
        "    if not optional_custom_model_id:\n",
        "        raise ValueError(\n",
        "            \"Please provide a optional_custom_model_id when 'other custom model' is selected.\"\n",
        "        )\n",
        "    model = (\n",
        "        f\"projects/{PROJECT_ID}/locations/{LOCATION}/models/{optional_custom_model_id}\"\n",
        "    )\n",
        "\n",
        "\n",
        "print(f\"Using model: {model}\")"
      ]
    },
    {
      "cell_type": "markdown",
      "id": "27ebf500",
      "metadata": {
        "id": "5c668aa7"
      },
      "source": [
        "Let's prepare the prediction set as multimodal dataset and start a Batch Prediction job with the specified model:"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "id": "52073a23",
      "metadata": {
        "id": "9daaf0ce"
      },
      "outputs": [],
      "source": [
        "# Create multimodal dataset for the prediction set and attach read config\n",
        "flowers_prediction = client.datasets.create_from_pandas(dataframe=prediction_set)\n",
        "\n",
        "flowers_prediction.set_read_config(read_config=read_config)\n",
        "flowers_prediction = client.datasets.update_multimodal_dataset(\n",
        "    multimodal_dataset=flowers_prediction\n",
        ")\n",
        "\n",
        "# Assemble the dataset with a read config and get the assembled BigQuery table id\n",
        "prediction_table_id, _ = client.datasets.assemble(\n",
        "    name=flowers_prediction.name,\n",
        "    gemini_request_read_config=read_config,\n",
        ")"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "id": "13617c2c",
      "metadata": {
        "id": "12b841c6"
      },
      "outputs": [],
      "source": [
        "# @title Run a Batch Prediction Job\n",
        "\n",
        "job = genai_client.batches.create(\n",
        "    # use the model selected above\n",
        "    model=model,\n",
        "    # Use the assembled BigQuery table as source\n",
        "    src=f\"bq://{prediction_table_id}\",\n",
        ")\n",
        "\n",
        "\n",
        "completed_states = {\n",
        "    JobState.JOB_STATE_SUCCEEDED,\n",
        "    JobState.JOB_STATE_FAILED,\n",
        "    JobState.JOB_STATE_CANCELLED,\n",
        "    JobState.JOB_STATE_PAUSED,\n",
        "}\n",
        "\n",
        "while job.state not in completed_states:\n",
        "    time.sleep(30)\n",
        "    job = genai_client.batches.get(name=job.name)\n",
        "    print(f\"Job state: {job.state}\")"
      ]
    }
  ],
  "metadata": {
    "colab": {
      "name": "intro_vertex_ai_multimodal_dataset.ipynb",
      "toc_visible": true
    },
    "kernelspec": {
      "display_name": "Python 3",
      "name": "python3"
    }
  },
  "nbformat": 4,
  "nbformat_minor": 0
}
