Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 0 additions & 2 deletions k8s/model-api/base/deployment.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -32,8 +32,6 @@ spec:
key: secret-access-key
- name: AWS_DEFAULT_REGION
value: ap-northeast-2
- name: CHURN_MODEL_VERSION # 모델 버전 추가
value: "v5"
resources:
requests:
cpu: "250m"
Expand Down
12 changes: 6 additions & 6 deletions services/model-api/app/models/loader.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,31 +6,31 @@
_model_cache: dict = {}

# mlflow에서 mnist 모델 로드하는 함수
def load_model(model_name: str, version: str):
cache_key = f"{model_name}_{version}"
def load_model(model_name: str):
cache_key = model_name
if cache_key in _model_cache:
return _model_cache[cache_key]

mlflow_uri = os.getenv("MLFLOW_TRACKING_URI", "http://mlflow:5000") # 환경변수에서 mlflow tracking uri 가져오기, 없으면 기본값으로 http://mlflow:5000 사용
mlflow.set_tracking_uri(mlflow_uri)

model_uri = f"models:/{model_name}-{version}/latest"
model_uri = f"models:/{model_name}/latest"
model = mlflow.pytorch.load_model(model_uri)
model.eval()

_model_cache[cache_key] = model
return model

# mlflow에서 고객 이탈 예측 모델 로드하는 함수
def load_churn_model(model_name: str, version: str):
cache_key = f"{model_name}_{version}_sklearn"
def load_churn_model(model_name: str):
cache_key = f"{model_name}_sklearn"
if cache_key in _model_cache:
return _model_cache[cache_key]

mlflow_uri = os.getenv("MLFLOW_TRACKING_URI", "http://mlflow:5000")
mlflow.set_tracking_uri(mlflow_uri)

model_uri = f"models:/{model_name}-{version}/latest"
model_uri = f"models:/{model_name}/latest"
model = mlflow.sklearn.load_model(model_uri)

_model_cache[cache_key] = model
Expand Down
4 changes: 1 addition & 3 deletions services/model-api/app/services/churn_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,6 @@
from prometheus_client import Counter, Histogram

MODEL_NAME = "churn"
MODEL_VERSION = os.getenv("CHURN_MODEL_VERSION", "v5")

FEATURES = [
'account_age_months',
Expand Down Expand Up @@ -36,7 +35,7 @@ def predict(data: dict) -> dict:
if missing:
raise ValueError(f"누락된 피처: {missing}")

model = load_churn_model(MODEL_NAME, MODEL_VERSION)
model = load_churn_model(MODEL_NAME)

x = [[data[f] for f in FEATURES]]

Expand All @@ -58,7 +57,6 @@ def predict(data: dict) -> dict:
},
"metadata": {
"model": MODEL_NAME,
"version": MODEL_VERSION,
"inference_time_ms": time_ms,
"timestamp": datetime.now(timezone.utc).isoformat(),
},
Expand Down
6 changes: 2 additions & 4 deletions services/model-api/app/services/mnist_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,8 +7,7 @@


MODEL_NAME = "mnist"
MODEL_VERSION = "v1"


prediction_counter = Counter(
"mnist_predictions_total",
"예측 횟수",
Expand All @@ -26,7 +25,7 @@ def predict(data: dict) -> dict:
if pixels is None or len(pixels) != 784:
raise ValueError(f"pixels 필드에 784개의 값이 필요합니다. 길이 오류: {len(pixels)}")

model = load_model(MODEL_NAME, MODEL_VERSION)
model = load_model(MODEL_NAME)

start = time.time()
with torch.no_grad():
Expand All @@ -48,7 +47,6 @@ def predict(data: dict) -> dict:
},
"metadata": {
"model": MODEL_NAME,
"version": MODEL_VERSION,
"inference_time_ms": time_ms,
"timestamp": datetime.now(timezone.utc).isoformat(),
},
Expand Down
2 changes: 1 addition & 1 deletion services/model-api/training/churn/train.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,7 +49,7 @@ def main(args):
mlflow.sklearn.log_model(
model,
artifact_path="model",
registered_model_name=f"churn-{args.version}",
registered_model_name="churn",
)

print(f"valid_f1 : {val_metrics['f1']:.4f}")
Expand Down
Loading