refactor: Simplify Pic2Word notebook steps (#2751)

* refactor: Simplify Pic2Word notebook steps

* fix: linter

* fix: Fix inference issue
This commit is contained in:
jismailyan-google
2024-02-28 22:05:33 +00:00
committed by GitHub
parent b4b20c3995
commit 9db64f74a2
3 changed files with 145 additions and 304 deletions
@@ -6,6 +6,7 @@ ENV infer_port=7080
ENV mng_port=7081
ENV model_name="pic2word"
ENV PATH="/home/model-server/:${PATH}"
ENV PYTHONPATH="$PYTHONPATH:/home/model-server/composed_image_retrieval:/home/model-server/composed_image_retrieval/src:/home/model-server"
# Copy license.
RUN apt-get update && apt-get install -y --no-install-recommends \
@@ -84,6 +85,7 @@ RUN pip uninstall dataclasses -y
# Copy model artifacts.
COPY model_oss/pic2word/handler.py /home/model-server/handler.py
COPY model_oss/util/ /home/model-server/util/
# Create torchserve configuration file.
RUN echo \
@@ -2,7 +2,7 @@
from argparse import Namespace # pylint: disable=g-importing-member
import os
from typing import Any
from typing import Any, List
from absl import logging
from data import CustomFolder
@@ -25,7 +25,7 @@ _COCO_DATASET_NAME = "coco"
_MODEL_NAME = "ViT-L/14"
_LOCAL_QUERY_PATH = "./query/"
_IMAGE_OUTPUT_LOCAL_DIR = "demo_out/images"
_OUTPUT_LOCAL_DIR = "/demo_out/"
_OUTPUT_LOCAL_DIR = "./demo_out/"
_DATA_DIR = "data"
_CHECKPOINT_DIR = "checkpoint/pic2word_model.pt"
_REQUEST_PROMPTS = "prompts"
@@ -33,6 +33,7 @@ _REQUEST_OUTPUT_STORAGE_DIR = "output_storage_dir"
_REQUEST_IMAGE_PATH = "image_path"
_REQUEST_IMAGE_FILE_NAME = "image_file_name"
_RESPONSE_MSG = "Successfully retrieved images."
_PICKLE_DIR_PATH = "gs://pic2word-bucket/pickle/"
class ModelHandler(BaseHandler):
@@ -49,6 +50,8 @@ class ModelHandler(BaseHandler):
def initialize(self, context: Any):
"""Initialize."""
logging.info("Initializing pic2word.")
# Download pickle file for COCO
fileutils.download_gcs_dir_to_local(_PICKLE_DIR_PATH, "./data")
# Download COCO dataset. The model looks for this folder specifically
# during image retrieval to generate a response for each request.
@@ -157,11 +160,11 @@ class ModelHandler(BaseHandler):
_IMAGE_OUTPUT_LOCAL_DIR, self.output_storage_dir
)
def handle(self, data: Any, context: Any) -> str: # pylint: disable=unused-argument
def handle(self, data: Any, context: Any) -> List[str]: # pylint: disable=unused-argument
"""Runs preprocess, inference, and post-processing."""
logging.info("Received Pic2Word inference request")
model_input = self.preprocess(data)
self.inference(model_input)
self.postprocess()
logging.info("Done handling input.")
return _RESPONSE_MSG
return [_RESPONSE_MSG]