varundevmishra09 commited on
Commit
85f8e77
·
verified ·
1 Parent(s): aaa7bd4

Create app.py

Browse files
Files changed (1) hide show
  1. app.py +126 -0
app.py ADDED
@@ -0,0 +1,126 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # app.py — Single-file Transit Tracker API for Hugging Face (No extra files required)
2
+ from fastapi import FastAPI, HTTPException
3
+ from pydantic import BaseModel
4
+ from typing import Dict, Any
5
+ import numpy as np
6
+ import joblib
7
+ import os
8
+ import math
9
+ import time
10
+
11
+ app = FastAPI(title="Transit Tracker API - Single File Version")
12
+
13
+ # -----------------------------
14
+ # In-memory Vehicle Locations
15
+ # -----------------------------
16
+ VEHICLES: Dict[str, Dict[str, Any]] = {}
17
+
18
+ # -----------------------------
19
+ # Train a lightweight ETA model
20
+ # -----------------------------
21
+ def train_and_save_model():
22
+ from sklearn.ensemble import RandomForestRegressor
23
+
24
+ np.random.seed(42)
25
+ N = 3000
26
+
27
+ dist = np.random.exponential(400, N)
28
+ speed = np.random.uniform(5, 15, N)
29
+ tod = np.random.uniform(0, 86400, N)
30
+ dow = np.random.randint(0, 7, N)
31
+
32
+ tod_sin = np.sin(2*np.pi * tod/86400)
33
+ tod_cos = np.cos(2*np.pi * tod/86400)
34
+
35
+ eta = dist/(speed*0.277) + np.random.randint(10, 120, N)
36
+
37
+ X = np.column_stack([dist, speed, tod_sin, tod_cos, dow])
38
+ y = eta
39
+
40
+ model = RandomForestRegressor(n_estimators=30)
41
+ model.fit(X, y)
42
+ joblib.dump(model, "model.joblib")
43
+ print("MODEL TRAINED AND SAVED ✔")
44
+
45
+ # -----------------------------
46
+ # Load or Train Model
47
+ # -----------------------------
48
+ def load_model():
49
+ if not os.path.exists("model.joblib"):
50
+ train_and_save_model()
51
+ return joblib.load("model.joblib")
52
+
53
+ MODEL = load_model()
54
+
55
+ # -----------------------------
56
+ # JSON Schemas
57
+ # -----------------------------
58
+ class VehicleUpdate(BaseModel):
59
+ vehicle_id: str
60
+ lat: float
61
+ lon: float
62
+ speed: float = 0.0
63
+
64
+ class PredictRequest(BaseModel):
65
+ distance_to_stop_m: float
66
+ speed_mps: float = 4.0
67
+
68
+
69
+ # -----------------------------
70
+ # Utils
71
+ # -----------------------------
72
+ def haversine(lat1, lon1, lat2, lon2):
73
+ R = 6371000
74
+ p1, p2 = math.radians(lat1), math.radians(lat2)
75
+ dphi = math.radians(lat2 - lat1)
76
+ dl = math.radians(lon2 - lon1)
77
+ a = math.sin(dphi / 2) ** 2 + math.cos(p1) * math.cos(p2) * math.sin(dl / 2) ** 2
78
+ return R * 2 * math.atan2(math.sqrt(a), math.sqrt(1 - a))
79
+
80
+
81
+ # -----------------------------
82
+ # ROOT
83
+ # -----------------------------
84
+ @app.get("/")
85
+ def root():
86
+ return {"message": "Transit Tracker API Running", "vehicles": len(VEHICLES)}
87
+
88
+
89
+ # -----------------------------
90
+ # LIVE VEHICLE LOCATION APIs
91
+ # -----------------------------
92
+ @app.post("/vehicle/update")
93
+ def update_vehicle(v: VehicleUpdate):
94
+ VEHICLES[v.vehicle_id] = {
95
+ "lat": v.lat,
96
+ "lon": v.lon,
97
+ "speed": v.speed,
98
+ "timestamp": time.time()
99
+ }
100
+ return {"success": True, "vehicle": VEHICLES[v.vehicle_id]}
101
+
102
+
103
+ @app.get("/vehicles")
104
+ def get_all_vehicles():
105
+ return VEHICLES
106
+
107
+
108
+ # -----------------------------
109
+ # ETA PREDICTION
110
+ # -----------------------------
111
+ @app.post("/predict")
112
+ def predict_eta(req: PredictRequest):
113
+ dist = req.distance_to_stop_m
114
+ if dist < 0:
115
+ raise HTTPException(400, "distance cannot be negative")
116
+
117
+ t = time.gmtime()
118
+ seconds = t.tm_hour * 3600 + t.tm_min * 60 + t.tm_sec
119
+ tod_sin = np.sin(2*np.pi * seconds/86400)
120
+ tod_cos = np.cos(2*np.pi * seconds/86400)
121
+ dow = t.tm_wday
122
+
123
+ X = np.array([[dist, req.speed_mps, tod_sin, tod_cos, dow]])
124
+ eta_sec = float(MODEL.predict(X)[0])
125
+
126
+ return {"eta_seconds": eta_sec, "eta_minutes": eta_sec / 60}