Support for cpu/gpu version

This commit is contained in:
Justin Güse
2024-07-12 11:07:53 +00:00
parent eb0f28b3f6
commit a8901d2431
7 changed files with 760 additions and 631 deletions
+9 -20
View File
@@ -1,25 +1,10 @@
# This project can be installed with `python3 -m pip install -e .` from the main directory.
[project]
name = "timesfm-jax"
[tool.poetry]
name = "timesfm"
packages = [
{ include = "*", from = "src" },
]
version = "0.0.1"
description = "Open weights time-series foundation model from Google Research."
version = "1.0.1"
dependencies = [
"jax==0.4.26",
"paxml==1.4.0",
"praxis==1.4.0",
"jaxlib==0.4.26",
"numpy==1.26.4",
"pandas==2.1.4",
"einshape==1.0.0",
"utilsforecast==0.1.10",
"huggingface_hub[cli]==0.23.0",
"scikit-learn==1.5.1",
]
version = "0.0.12"
authors = [
"Rajat Sen <senrajat@google.com>",
"Yichen Zhou <yichenzhou@google.com>",
@@ -49,9 +34,13 @@ einshape = ">=1.0.0"
numpy = ">=1.26.4"
pandas = ">=2.1.4"
paxml = "1.4.0"
jax = "0.4.26"
jaxlib = "0.4.26"
utilsforecast = "^0.1.10"
jax = {version = "0.4.26", extras = ["cuda12"]}
# unfortunately we have to solve it using https://stackoverflow.com/questions/72037181/how-to-add-optional-dependencies-of-a-library-as-extra-in-poetry-and-pyproject
Jax = {version = "0.4.26", extras = ["cpu"], optional = true}
[tool.poetry.extras]
cpu = ['Jax']
[build-system]