Add Notebook

This commit is contained in:
vowelparrot committed 2023-04-17 23:04:41 -07:00
1 parent 46e690200c
commit 80937a1330
3 files changed
+192

No files matched your search

@@ -0,0 +1,98 @@
{
"cells": [
{
"cell_type": "markdown",
"metadata": {},
"source": [
"# HuggingFace Image Generation Tool\n",
"\n",
"Agents can generate models using HuggingFace Inference Enpdoints via the HuggingFace Image Generation Tool.\n",
"\n",
"To access the endpoint, you will have to copy your API Key from your [HuggingFace console](https://huggingface.co/settings/tokens)."
]
},
{
"cell_type": "code",
"execution_count": 1,
"metadata": {
"tags": []
},
"outputs": [],
"source": [
"from langchain.tools.huggingface.image_generation import HuggingFaceImageGenerationTool"
]
},
{
"cell_type": "code",
"execution_count": 2,
"metadata": {
"tags": []
},
"outputs": [],
"source": [
"api_key = \"\" # Copy a token with read permissions from https://huggingface.co/settings/tokens\n",
"image_generation_tool = HuggingFaceImageGenerationTool.from_api_key(api_key, \"CompVis/stable-diffusion-v1-4\")"
]
},
{
"cell_type": "code",
"execution_count": 3,
"metadata": {
"tags": []
},
"outputs": [
{
"data": {
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAgAAAAIACAIAAAB7GkOtAAEAAElEQVR4nKT9WcxtW3Yehn1jjDnXWnvv//9Pd89t6rbVsRqSxSJZJItVZJE0TauJ4kYWEwWJQMd2AgixHxLlIUD8HCABAgRGgjzEgWNAkBAEomw1UCQakCVZpEiTFhsVi6oqFm9V3f6e7m92s+YcTR7G2vuckmj4IRsH5/73/HuvvZo5R/ONb3yDvvl/+EUiYhAzEwVxkLCTkaDFvDqrNNB2f9WpT9PEDEHtvXsYmJw8IrhQrbU3AwBnAOwFAPvyBgslolorgs0sIpjZLABEgJmZGUBEuDtJJSJ39TAS1IGIYD4HqRSUwixhEYEOgIjQg7kQEYCwcPf82bvXWrWZmY3jGBH5w9wPABDGTIUY5GTm7gVERGzi7oULM3vD3PbC4a4cIKLo3nsv4LFUIQkzmy0iSlBEwN3VSBkeEogItvxbKBAGMg8Nbd1nba1Z1zDlZlIouCp87hFOg9SJh8PlgdW9m6o29dm1mTaKWA9zc1gZ5azgwlwORgeJty4f2r3yZNw/WW2/sW13PonVCyPdxYfz/INfuLfdXe/3LQjPv3Bxtb26c+ei6fzg4XxxG2dn47RZrzfDbrdznYciIrJe3aqri91Ov/vuw3feenR9AxIAdLG+ePvty3mHu3fBjvNbE/F0OOj28nD1QJ+8i+sPIXusDZuQw/v2g298Yv/hlvdxZ33x/pvvvvLyCwHd76+j+Wa16nNrh/ls2sytXdw6u9ptjWCMq8ONUgzrcbBytpOLcmbhPRyTGKJxTKvVoe0jAuEUkICAlj8kMHd3AMEUEbmutvtdKUVEmDkicvEE0aH1abUiov1+H+7DMBBRa02IAeTbCuUiZSLyNjNzKYWICBERHAAw1UJEQkFEFMfF7DGUOh8OgK9WKxFRVUYMw7DdbkWEGUR03AJuZlMZVRVwEaGAuxOHiIQ5ETE8V7iIEFFeSF4jgCDKoxHRzW67Xq+llt67meX71T1kUDcAIqLq8zx7RK0V4IhQj4gIykuHg+e5SxnA1N084ELdopnOhhYW4BB2IjULQmFBRCHOc3M1AEMppRRtnZmFmJkFy3OJiKGWUkpEzPMM81qrFHJ3Or4FAFGcrou9qCoAFgBwbwU0FBrrQNbCrSAGJiGuhLyQICIOAJxH4yCKKsW8eVeQ51Nwd1crXIUoqARgRopACJjmec4TMFNTJYqpDsMwzV2JJC0YAHc3h7sHE5iIxCM0EBEgIiJhjojwZWkRUQS5e+89IgBm5lylBAF5aGdZnmxELIc4fp2Gq6qqujuWW1QjggKnF3nk4wbARMBicgVERMUdgHngeIuDwoMxVPHurWkRJhLhAKDqpVJh1og8D3cnRtrc0+u4hP7Yf/z/67XcOwLAgH/P6o8IXxbNv3A+//JBTr8+3SkzYy4U4e7qSpDQcHde3srw5f2BACjMXC0iyJevdVXXGEhACER4GEAGd4fHfHNABCnCPNThQUTCxcLYJYjh5Kaq6rMbGQd5wILU0TwMhFJLYQswlrsfjN7tAN+qrp9fX9X5Wq9u3GSN2/en268+V+7S1N45v5i0b7uCCCI43wy3bm+AldQHFxcXUhERoeo6m3WuxERqc7t+/OT6cNjfjBONq0F4NItpGO7cgW7wwv3zwnTr1q07955rc9/d9KsHNx/ef/Thd/f7B5AZa6M7m/GqvysrKmM94PHtN4Z9PNFZd21LTuTEHDKKcwTTdj+3pnUYah0CrGREHO4K6hSK6BQC6jAzV1U3iwiKIICCmEiIGGRm7g5zQ8AWQ2zhZ2dnz64B99zvVmtVVSIahiHczQxArZWfWb1p0M2MiMpQ5dlfEQkzETXtnCfCke9fDBxqrTXC0ie5u7vloSICoNzMZgb4Yobc82cGRQSO7oqIEMteyzUPID/u7nH8p/zvMAzu7r0DKKXkVn3GpNIzx4yIYCZmJiEzU7fcRIXZSmERx8nT4Pi94OAABzOI+Ljv4plXfkXeujRAeXXmT3fo6Z3MTEsYCmbufY6nFx75ccLTG0tMzBCRIlIqWzeJYGZGnM6UQI4IINwB5LakNKOh5gZ3kBMJcDwHIqYSzAEOgAMeALBer/PkiWiolRkZFtRazUJVc+UwM0utte7mw3IGRCKFmSOvJR/W8dkRkXu+R9w9YnHqZoaIgK2GCnI/PlZmllJKKbvdDsgFfvSOzMysulzpaX2eFsz3PPTjzS+1DIHTyiBmIgIYRAyDdqcBRMKMcJi6sYU7CMzs5EQUjnB6Nh7577K+z66M0789Pdl/6c0cRLH4kQgKB4QjLBCnPZg7PLdKLnFJP7dcpOPp1WHxjQHkskfk+j8+Di5cALh6hC+bLUBO4QGi8FyE5A7vjkAEIRAeoW7qruHF4YB5mLsGmcMAD2aGhptZU+turZsZeUgLFUKFB0LhHaqm7gONcNKg2bG3UAKGUqS4Nu+hfdbCKDgEbmDbon2gq7qfJ1u9SC+ex6ufffn8xVt0pi9M5fzuMA1t7lUKr85W7tNzz99WbXXsZ2dnHm2/34O1lhgKb9ZSayXheTaRvlnz+flmtb5grm1W1yBfeejd2ysR2qz94uI6zO5c8At3xxefv3v9+qFfe7Hp8TtXZ3Lxz3/33RrgwM0W9y/Wb7+5e+2l+/F4JCUK+CFIuWlooFCBlTARq2uIhcHpEK0RtoUM7O5FYEZdlXu3rhwQEIMEwRTMy9P0oAD5YoXhju5R43vspjvMQt2nQeZ5FubVauXurTUKjHXIrZ6Z33EFLpFUhvZLJE5MTMzsysHhBCZhIkQAxoHWmoiARC0oVEQ8ovU+DNXcCVKkhvXeOxGN4xhH4wBwUICJwAQh9gx+ImMf99zMmdDkv+cWABOAWmvv3T1qrRliL1tADQQIPWsg/sWwySPcQRQUaaSQftQ8iNNICYu4BSEYDDIGBTPnWS7fJYs1gbsPw7gYf/eIIJCAmDl94XIaDCfw8oCWCwTARPmHjumShxLxMwaNuioXpAMIU49wYhEWyUviCAOBiAmOgPXurhFOHAwiZlocPaXVXr785MmkmjVTBSBFmCXceu9BHPHUhgB8tLPi4e4mIkyUd8bdI92w4xTRZwZwvA9yuhsEAQkR4ug482rlWf/9L5nNDAgonhr6jFcyyjklUqc7XEQk13PAEJSeFYSwyIsh1HBbQh1G7z1PXCQyYAkgkzKAgePOWzYME2la+Ixlvtf6//F2/3v/XlbnM5swAgEEswgoyNNMn17HNf3H5CXfc6ci7yIAplD34GAQImAaEc4BODGTmUaQAHAiZ8rNpkEAOcEBC9cwDVdv3uFBTmHuamQgDzIaqAacXMPgXU3V1djIGgcHdXficA4FOSFw6MrBFtSDGujgHuY1rEbxbvuuZrO6HQr2Y+xX/t3dtr6Es1fqC599qW/8pU/do7VYcYNa39fapk0Zx6reWztYK1Klchfah7UifbUaV8MExOWTm2lCHdYRqNzH0Yv0sc7TWHSgPne+PwjV9ZoZAO/m3UNEZxrKMN67X+/cHQbaTGVz+eHmbLpdn7sWlxLSdvPdi7v8tf3u/ceHriVkfxMWvqqbQlxlrLKeD6YdmSEzi1QuEXOJXsidgqSQONQdppp3aQkKiTyip3/m4uS+RACUz0iC9vs9F6lSggAPCyfhoUhuvDSjESEiGTD03pn5BP7k1ooIszSznvvKKbq5uAeB4mnakBYEhO3NdhzHwmJmpZRxLM5sZk8XNpD7cwFCj/v2FKhxIioeaUAXi++ep5Rbz57dNR44moygxdBkiEpc3BzCp62RkWm+X1XtaMvyDarqYIRZuJn5EhcxES3nR8QgBRhLqFZZLMzd+fjCMXBmZpLlNp6cq2rHMRbOU834uJRy2st89ENPXQCVI2jm3T0c7lajEDEBERZuzsQRjOKhEUGUmAFFUISNZQwUwIFgImYQCQVaawYwOMCGSKvtYM4HD2a4B8g9Ah7uakQiIgmj5fu7GxcxzQ/7giuYqWoROT32Zw1ULgB55tEQEYvA7XTVzzrpYRgiAvEUBTlGwGnVnh4kA9l/KfJePlXMwj3MPO88OxwBRwkQZChVWHqfLVCkDKXofPAwyoyMCMGBhPULEB6eccgp7+AFD8yFZQFDJAjreYZ5Ms/YZVA+lqP1z7fQktgBiAT3OJBRAVMkgvl0xTAjjLCcBh9BGnJL/wSKDJTgwQEP1t6NDQTvERoRMI+IqIXJKTxxH6ZjLmkWDBIHBdwiLOAAWLsjgp0zP6AgihDwvN+7ehzMu4dGCSEWJvEOuJiHBdwBCDsDcNMgDnAILFi9d8NsvkEJCJViA+1KXJe2Xdv+TLcb3PsY7n/29oufuXdpl3KnRaXVaL31Nj8eR9RpXK/rdnvgwYBtP9hhdxN+4xEiWA0TBgFi3qMWEO2L8GYj683AVIl9Nz9kLqXYvXsyjbUUYaA3cz2EA7EX2buXKMVo38pNfYF4wMe+eHfeNyIZy20KlI+8+Ogb/fKdufr04NuPH33XKfZxrWu5KGydiEFAellmZTLmQiBBOAczRELcmToTiVAAcLhHBJbwohK5IZ5Cmrk8Yj5sxzqUOgZBWwdIShnH8erJ5TiOzDwfekTUMhLQe+dgcvLj2qPMRRnmnrk4RFhEEGZm6hQJ4DAzjBZzzAENSAAeHs4IDwIJCXuujjQRQUHiILNghGUBg0iIIhBkYAJC3dxDCBShZgAKLeuRKF0U5fYAMDcFqEgJp9lM1QAUXhJ0CoR73qUFlQaZmaoxcxkGETEzU/NwkuP9FBALSAhuDo6nezi3FaXRIiYGM1cRZjYgHeqyN5fnlrlakHB6qVMh8PhzogqZgT1N4mFwdykn6Ayq6uZLXQHIxxOgCPIIPrqH43qIzG/GzQhyeH4iQCEkRGTmtNisk0kqAIFFKnOpcAM5gZiogg6HJiCQZJ2BgGABTh5L4hnXlY7tlAGkg8yMKlEy5pI/uztCiWMLine truncated
"text/plain": [
"<IPython.core.display.Image object>"
]
},
"execution_count": 3,
"metadata": {},
"output_type": "execute_result"
}
],
"source": [
"from pathlib import Path\n",
"from tempfile import TemporaryDirectory\n",
"\n",
"from IPython.display import Image\n",
"\n",
"\n",
"out_path = \"parrot.png\"\n",
"image_generation_tool.run({\"prompt\": \"a painting of a parrot playing a saxaphone\", \"out_path\": str(out_path)})\n",
"# Show the generated image\n",
"Image(out_path)"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": []
}
],
"metadata": {
"kernelspec": {
"display_name": "Python 3 (ipykernel)",
"language": "python",
"name": "python3"
},
"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.11.2"
}
},
"nbformat": 4,
"nbformat_minor": 4
}
Whitespace-only changes.
@@ -0,0 +1,94 @@
"""HuggingFace Image Generation Wrapper."""
from __future__ import annotations
import io
from typing import Optional
from PIL import Image
from pydantic import BaseModel, Field, root_validator
from requests import Response
from tenacity import retry, retry_if_result, stop_after_attempt, wait_fixed
from langchain.requests import Requests
from langchain.tools.base import BaseTool
def is_503_error(response: Response) -> bool:
return response.status_code == 503
class RunArgsSchema(BaseModel):
"""Schema for the RunArgs."""
prompt: str = Field(..., description="Prompt to generate an image.")
out_path: str = Field(..., description="Path to write the generated image to.")
DEFAULT_INFERENCE_URL = "https://api-inference.huggingface.co/models/"
class HuggingFaceImageGenerationTool(BaseTool):
"""Image Generation Wrapper."""
name = "huggingface_image_generation"
description = (
"Generate an image using a valid image generation model from HuggingFace's API."
)
requests_wrapper: Requests
model_id: str
"""The id of the model to use, such as 'CompVis/stable-diffusion-v1-4'."""
"""Requests wrapper to use containing the authorization headers."""
url_base: str = DEFAULT_INFERENCE_URL
@root_validator
def _validate_authorization_present(cls, values: dict) -> dict:
requests: Requests = values["requests_wrapper"]
headers = requests.headers or {}
if headers.get("Authorization") is None:
raise ValueError(
"Error: Authorization token required for the requests wrapper of"
" the HuggingFaceImageGenerationTool."
)
return values
@retry(
stop=stop_after_attempt(5),
wait=wait_fixed(20),
retry=retry_if_result(is_503_error),
)
def _request_huggingface_image(
self,
prompt: str,
) -> Response:
"""Generate an image using Huggingface's API."""
api_url = self.url_base + self.model_id
response = self.requests_wrapper.post(
api_url,
data={
"inputs": prompt,
},
)
return response
def _run(self, prompt: str, out_path: str) -> str:
"""Generate an image using Stable Diffusion using HuggingFace's API."""
response = self._request_huggingface_image(prompt=prompt)
if response.status_code == 200:
image = Image.open(io.BytesIO(response.content))
image.save(out_path)
return f"Saved to disk: {out_path}"
else:
return f"Failed to generate image. Error: {str(response.content)}"
async def _arun(self, prompt: str, out_path: str, model_id: str) -> str:
raise NotImplementedError
@classmethod
def from_api_key(
cls, huggingface_api_key: str, model_id: str, url_base: Optional[str] = None
) -> HuggingFaceImageGenerationTool:
"""Create a HuggingFaceImageGenerationTool from an API key."""
requests = Requests(headers={"Authorization": f"Bearer {huggingface_api_key}"})
url_base = url_base or DEFAULT_INFERENCE_URL
return cls(requests_wrapper=requests, model_id=model_id, url_base=url_base)