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로 노출하여 외부에서 슬롯 번호 계산 검증을 볼 수 있기 하였다
'Project > SSAFY2학기 특화 PJT' 카테고리의 다른 글
| [BE, AI] 언어 선택 이유 (0) | 2026.03.13 |
|---|---|
| [BE] odsay.rs ( ODSAY API 연동 ) (0) | 2026.03.13 |
| [AI] train.py ( 모델 훈련 ) (0) | 2026.03.13 |
| [AI] preprocess.py ( 데이터 전처리 ) (0) | 2026.03.13 |
| 팀 프로젝트 정리 사이트 (0) | 2026.03.03 |
