Spaces:
Running
Running
Upload 2 files
Browse files- schemas.py +2 -0
- service.py +9 -6
schemas.py
CHANGED
|
@@ -14,6 +14,8 @@ class HealthResponse(BaseModel):
|
|
| 14 |
ready: bool
|
| 15 |
max_context_length: int
|
| 16 |
max_horizon_step: int
|
|
|
|
|
|
|
| 17 |
|
| 18 |
|
| 19 |
class PredictRequest(BaseModel):
|
|
|
|
| 14 |
ready: bool
|
| 15 |
max_context_length: int
|
| 16 |
max_horizon_step: int
|
| 17 |
+
patch_length: int
|
| 18 |
+
runtime_revision: str
|
| 19 |
|
| 20 |
|
| 21 |
class PredictRequest(BaseModel):
|
service.py
CHANGED
|
@@ -22,8 +22,9 @@ class TimesFmService:
|
|
| 22 |
"TIMESFM_MODEL_NAME",
|
| 23 |
"google/timesfm-2.5-200m-transformers",
|
| 24 |
)
|
| 25 |
-
self.backend = os.getenv("TIMESFM_BACKEND", "hf_cpu").strip() or "hf_cpu"
|
| 26 |
-
self.device = "cpu"
|
|
|
|
| 27 |
self.max_context_length = int(os.getenv("TIMESFM_MAX_CONTEXT_LENGTH", "512"))
|
| 28 |
self.max_horizon_step = int(os.getenv("TIMESFM_MAX_HORIZON_STEP", "288"))
|
| 29 |
self.patch_length = int(os.getenv("TIMESFM_PATCH_LENGTH", "32"))
|
|
@@ -46,10 +47,12 @@ class TimesFmService:
|
|
| 46 |
model_id=self.model_id,
|
| 47 |
backend=self.backend,
|
| 48 |
device=self.device,
|
| 49 |
-
ready=self.ready,
|
| 50 |
-
max_context_length=self.max_context_length,
|
| 51 |
-
max_horizon_step=self.max_horizon_step,
|
| 52 |
-
|
|
|
|
|
|
|
| 53 |
|
| 54 |
def predict(self, payload: PredictRequest) -> PredictResponse:
|
| 55 |
self._validate_request(payload)
|
|
|
|
| 22 |
"TIMESFM_MODEL_NAME",
|
| 23 |
"google/timesfm-2.5-200m-transformers",
|
| 24 |
)
|
| 25 |
+
self.backend = os.getenv("TIMESFM_BACKEND", "hf_cpu").strip() or "hf_cpu"
|
| 26 |
+
self.device = "cpu"
|
| 27 |
+
self.runtime_revision = os.getenv("TIMESFM_RUNTIME_REVISION", "timesfm-hf-patch-align-v1")
|
| 28 |
self.max_context_length = int(os.getenv("TIMESFM_MAX_CONTEXT_LENGTH", "512"))
|
| 29 |
self.max_horizon_step = int(os.getenv("TIMESFM_MAX_HORIZON_STEP", "288"))
|
| 30 |
self.patch_length = int(os.getenv("TIMESFM_PATCH_LENGTH", "32"))
|
|
|
|
| 47 |
model_id=self.model_id,
|
| 48 |
backend=self.backend,
|
| 49 |
device=self.device,
|
| 50 |
+
ready=self.ready,
|
| 51 |
+
max_context_length=self.max_context_length,
|
| 52 |
+
max_horizon_step=self.max_horizon_step,
|
| 53 |
+
patch_length=self.patch_length,
|
| 54 |
+
runtime_revision=self.runtime_revision,
|
| 55 |
+
)
|
| 56 |
|
| 57 |
def predict(self, payload: PredictRequest) -> PredictResponse:
|
| 58 |
self._validate_request(payload)
|