Taylor1998 commited on
Commit
ea1a58e
·
verified ·
1 Parent(s): 9fba3c2

Upload 2 files

Browse files
Files changed (2) hide show
  1. schemas.py +2 -0
  2. 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)