From 369bc0a70ed81565011ad95f331e5a3947b7c85f Mon Sep 17 00:00:00 2001 From: Alex Liao Date: Sat, 20 Dec 2025 11:52:32 -0800 Subject: [PATCH] Auto-detect device for easier getting started --- .gitignore | 1 + README.md | 2 +- examples/prediction_example.py | 2 +- model/kronos.py | 12 +++++++++++- requirements.txt | 2 +- 5 files changed, 15 insertions(+), 4 deletions(-) diff --git a/.gitignore b/.gitignore index 3054aca..68871c9 100644 --- a/.gitignore +++ b/.gitignore @@ -73,3 +73,4 @@ venv.bak/ *.temp temp/ tmp/ +.python-version diff --git a/README.md b/README.md index fb507d7..faffa15 100644 --- a/README.md +++ b/README.md @@ -118,7 +118,7 @@ Create an instance of `KronosPredictor`, passing the model, tokenizer, and desir ```python # Initialize the predictor -predictor = KronosPredictor(model, tokenizer, device="cuda:0", max_context=512) +predictor = KronosPredictor(model, tokenizer, max_context=512) ``` #### 3. Prepare Input Data diff --git a/examples/prediction_example.py b/examples/prediction_example.py index 0304848..880f22b 100644 --- a/examples/prediction_example.py +++ b/examples/prediction_example.py @@ -43,7 +43,7 @@ tokenizer = KronosTokenizer.from_pretrained("NeoQuasar/Kronos-Tokenizer-base") model = Kronos.from_pretrained("NeoQuasar/Kronos-small") # 2. Instantiate Predictor -predictor = KronosPredictor(model, tokenizer, device="cuda:0", max_context=512) +predictor = KronosPredictor(model, tokenizer, max_context=512) # 3. Prepare Data df = pd.read_csv("./data/XSHG_5min_600977.csv") diff --git a/model/kronos.py b/model/kronos.py index e7ebba6..4696014 100644 --- a/model/kronos.py +++ b/model/kronos.py @@ -481,7 +481,7 @@ def calc_time_stamps(x_timestamp): class KronosPredictor: - def __init__(self, model, tokenizer, device="cuda:0", max_context=512, clip=5): + def __init__(self, model, tokenizer, device=None, max_context=512, clip=5): self.tokenizer = tokenizer self.model = model self.max_context = max_context @@ -490,6 +490,16 @@ class KronosPredictor: self.vol_col = 'volume' self.amt_vol = 'amount' self.time_cols = ['minute', 'hour', 'weekday', 'day', 'month'] + + # Auto-detect device if not specified + if device is None: + if torch.cuda.is_available(): + device = "cuda:0" + elif hasattr(torch.backends, 'mps') and torch.backends.mps.is_available(): + device = "mps" + else: + device = "cpu" + self.device = device self.tokenizer = self.tokenizer.to(self.device) diff --git a/requirements.txt b/requirements.txt index d94a8d7..598a3de 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,6 +1,6 @@ numpy pandas -torch +torch>=2.0.0 einops==0.8.1 huggingface_hub==0.33.1