Auto-detect device for easier getting started

This commit is contained in:
Alex Liao
2025-12-20 11:52:32 -08:00
parent 391bfab26e
commit 369bc0a70e
5 changed files with 15 additions and 4 deletions
+1
View File
@@ -73,3 +73,4 @@ venv.bak/
*.temp
temp/
tmp/
.python-version
+1 -1
View File
@@ -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
+1 -1
View File
@@ -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")
+11 -1
View File
@@ -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)
+1 -1
View File
@@ -1,6 +1,6 @@
numpy
pandas
torch
torch>=2.0.0
einops==0.8.1
huggingface_hub==0.33.1