mirror of
https://github.com/langchain-ai/langchain.git
synced 2026-10-05 17:35:28 +03:00
Add Notebook
This commit is contained in:
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)
|
||||
Reference in new issue
Block a user