model = None
encoders = {}

@asynccontextmanager
async def lifespan(app: FastAPI):
    global model, encoders
    
    try:
        model = joblib.load("models/congestion_model.pkl")
        for col in ["요일구분", "호선", "출발역", "상하구분"]:
            encoders[col] = joblib.load(f"models/encoders/{col}_encoder.pkl")
        print("모델 & 인코더 로드 완료")
        
    except FileNotFoundError as e:
        print(f"모델 파일 없음 — train.py 먼저 실행하세요: {e}")
    yield

So what: 서버 시작 시 모델과 인코더를 한 번만 로드 후 전역 변수로 유지한다

So why: 요청마다 로드를 할 경우 매번 디스크에서 파일을 읽어야 하므로 응답 시간이 늘어난다. lifespan으로 서버 시작 시점에 한 번만 로드해두면 이후 요청은 메모리에서 바로 꺼내 쓸 수 있다.

 

def safe_encode(col: str, val: str) -> int:
    if val not in encoders[col].classes_:
        raise ValueError(
            f"'{val}'은 '{col}'에서 지원하지 않는 값입니다. "
            f"가능한 값: {list(encoders[col].classes_)}"
        )
    return int(encoders[col].transform([val])[0])

So what: 인코더가 모르는 값이 들어오면 에러 메세지와 함께 가능한 값 목록을 반환해준다.

So why: LabelEncoder는 학습 때 없던 값이 들어오면 내부에서 예외가 발생하는데 미리 검사해서 명확한 메세지를 주면 백엔드에서 400 Error를 받았을 때 원인을 바로 알 수 있다. 이를 통해 컬럼은 {}호선 형태와  수도권 {}호선 형태가 나눠져있다는 것을 찾아내서 수정을 할 수 있었다.

 

def _predict_one(req: PredictRequest) -> PredictResponse:
    features = np.array([[
        safe_encode("요일구분", req.day_type),
        safe_encode("호선",    req.line),
        req.station_no,
        safe_encode("출발역",  req.station_name),
        safe_encode("상하구분", req.direction),
        req.time_slot
    ]])
    pred = float(model.predict(features)[0])
    pred = max(0.0, round(pred, 1))

So what: 요청 파라미터를 train.py의 FEATURES 순서와 동일하게 배열로 만들어서 예측한다.

So why: LightGBM은 학습 때의 피처 순서 그대로 입력받아야 한다. 순서가 하나라도 다를 경우 이상한 값을 예측하게 된다. max(0,0 ...)은 모델이 음수를 예측하는 경우를 막기 위해서이다. 혼잡도는 0% 미만이 물리적으로 불가능하기 때문이다

 

@app.post("/predict/batch", response_model=List[PredictResponse])
def predict_batch(requests: List[PredictRequest]):
    """여러 구간 혼잡도 일괄 예측 (경로 정렬용)"""
    if not requests:
        return []
    try:
        return [_predict_one(req) for req in requests]
    except ValueError as e:
        raise HTTPException(status_code=400, detail=str(e))

So what: 여러 구간을 한 번의 요청으로 예측한다

So why: 백엔드에서 경로별로 지하철 구간이 여러 개일 때 /predict를 구간 수만큼 개별 호출할 경우 HTTP 연결 오버헤드가 구간 수만큼 발생한다. batch로 묶어서 한 번에 보내면 네트워크 왕복 1회로 줄일 수 있다.

 

@app.get("/health")
def health():
    return {"status": "ok", "model_loaded": model is not None}

@app.get("/slot")
def time_to_slot(hour: int, minute: int):
    slot = hour * 2 + (1 if minute >= 30 else 0)
    return {"시간슬롯": slot}

So what: /health는 서버 상태와 모델 로드 여부를 반환한다. /slot은 시각을 슬롯 번호로 변환해준다

So why: model_loaded를 따로  보내는 이유는 서버가 커져 있어도 모델 로드에 실패한 상태일 수 있기 때문에 모델 상태를 확인하고자 넣었다. /slot은 preprocess.py의 time_to_slot 함수와 동일한 로직을 API로 노출하여 외부에서 슬롯 번호 계산 검증을 볼 수 있기 하였다

+ Recent posts