From 049582d9d6db6fe2b89dd0598ed93821f77cade9 Mon Sep 17 00:00:00 2001 From: Cuong Tham Date: Sat, 24 Jun 2023 15:09:36 -0700 Subject: [PATCH] Week3 project for MLOps: From Models to Production --- week3/project/app/classifier.py | 19 +++++++++++++---- week3/project/app/server.py | 36 ++++++++++++++++++++++++++++++++- week3/project/e2e.sh | 11 ++++++++++ 3 files changed, 61 insertions(+), 5 deletions(-) create mode 100755 week3/project/e2e.sh diff --git a/week3/project/app/classifier.py b/week3/project/app/classifier.py index 0757ad5..b669b94 100644 --- a/week3/project/app/classifier.py +++ b/week3/project/app/classifier.py @@ -1,8 +1,6 @@ from typing import List - from loguru import logger import joblib - from sentence_transformers import SentenceTransformer from sklearn.base import BaseEstimator, TransformerMixin from sklearn.pipeline import Pipeline @@ -72,7 +70,12 @@ def predict_proba(self, model_input: dict) -> dict: ... } """ - return {} + X = [ + model_input['title'] + ' | ' + model_input['description'] + ] + probs = self.pipeline.predict_proba(X)[0] + + return dict(zip(self.classes, probs)) def predict_label(self, model_input: dict) -> str: """ @@ -83,4 +86,12 @@ def predict_label(self, model_input: dict) -> str: Output format: predicted label for the model input """ - return "" + probs = self.predict_proba(model_input) + highest_prob = 0 + label = "" + for k, v in probs.items(): + if v > highest_prob: + highest_prob = v + label = k + + return label diff --git a/week3/project/app/server.py b/week3/project/app/server.py index 36f998c..64a56f6 100644 --- a/week3/project/app/server.py +++ b/week3/project/app/server.py @@ -1,3 +1,7 @@ +import json +import math +import time +from datetime import datetime from fastapi import FastAPI from pydantic import BaseModel from loguru import logger @@ -19,6 +23,7 @@ class PredictResponse(BaseModel): MODEL_PATH = "../data/news_classifier.joblib" LOGS_OUTPUT_PATH = "../data/logs.out" +global_data = {} app = FastAPI() @@ -34,6 +39,11 @@ def startup_event(): Access to the model instance and log file will be needed in /predict endpoint, make sure you store them as global variables """ + + global_data['file_log_handler'] = open(LOGS_OUTPUT_PATH, 'w+') + global_data['cls'] = NewsCategoryClassifier(verbose=True) + global_data['cls'].load(MODEL_PATH) + logger.info("Setup completed") @@ -45,6 +55,10 @@ def shutdown_event(): 1. Make sure to flush the log file and close any file pointers to avoid corruption 2. Any other cleanups """ + + if global_data['file_log_handler'] is not None: + global_data['file_log_handler'].close() + global_data['file_log_handler'] = None logger.info("Shutting down application") @@ -65,7 +79,27 @@ def predict(request: PredictRequest): } 3. Construct an instance of `PredictResponse` and return """ - response = PredictResponse(scores={"label1": 0.9, "label2": 0.1}, label="label1") + + t0 = time.monotonic_ns() + X = { + 'source': request.source, + 'url': request.url, + 'title': request.title, + 'description': request.description + } + probs = global_data['cls'].predict_proba(X) + label = global_data['cls'].predict_label(X) + response = PredictResponse(scores=probs, label=label) + t1 = time.monotonic_ns() + + log = { + 'timestamp': datetime.now().strftime('%Y-%m-%d %H:%M:%S'), + 'request': request.dict(), + 'latency(ms)': (t1 - t0) / 1000000 + } + global_data['file_log_handler'].write(json.dumps(log) + '\n') + global_data['file_log_handler'].flush() + return response diff --git a/week3/project/e2e.sh b/week3/project/e2e.sh new file mode 100755 index 0000000..f7296be --- /dev/null +++ b/week3/project/e2e.sh @@ -0,0 +1,11 @@ +#!/bin/bash -e + +while IFS= read -r line; do + + curl --location 'http://localhost/predict' \ + --header 'accept: application/json' \ + --header 'Content-Type: application/json' \ + --data "$line" + echo "" + +done < data/requests.json