Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -246,7 +246,7 @@ All settings for data, training, and model paths are centralized in `finetune/co
* `backtest_result_path`: Directory for saving backtesting results.
* `pretrained_tokenizer_path` and `pretrained_predictor_path`: Paths to the pre-trained models you want to start from (can be local paths or Hugging Face model names).

You can also adjust other parameters like `instrument`, `train_time_range`, `epochs`, and `batch_size` to fit your specific task. If you don't use [Comet.ml](https://www.comet.com/), set `use_comet = False`.
You can also adjust other parameters like `instrument`, `train_time_range`, `epochs`, and `batch_size` to fit your specific task. [Comet.ml](https://www.comet.com/) logging is disabled by default. To enable it, install `comet_ml`, set `use_comet = True`, and configure the Comet credentials in `config.py`.

### Step 2: Prepare the Dataset

Expand Down
4 changes: 3 additions & 1 deletion finetune/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -72,7 +72,9 @@ def __init__(self):
# =================================================================
# Experiment Logging & Saving
# =================================================================
self.use_comet = True # Set to False if you don't want to use Comet ML
# Comet is optional. Set to True after installing ``comet_ml`` and
# configuring the credentials below.
self.use_comet = False
self.comet_config = {
# It is highly recommended to load secrets from environment variables
# for security purposes. Example: os.getenv("COMET_API_KEY")
Expand Down
22 changes: 7 additions & 15 deletions finetune/train_predictor.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,23 +7,22 @@
import torch
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
from torch.nn.parallel import DistributedDataParallel as DDP

import comet_ml
from torch.nn.parallel import DistributedDataParallel as DDP

# Ensure project root is in path
sys.path.append('../')
from config import Config
from dataset import QlibDataset
from model.kronos import KronosTokenizer, Kronos
# Import shared utilities
from utils.training_utils import (
from utils.training_utils import (
setup_ddp,
cleanup_ddp,
set_seed,
get_model_size,
format_time
)
format_time
)
from utils.experiment_logging import create_comet_experiment


def create_dataloaders(config: dict, rank: int, world_size: int):
Expand Down Expand Up @@ -197,15 +196,8 @@ def main(config: dict):
'world_size': world_size,
}
if config['use_comet']:
comet_logger = comet_ml.Experiment(
api_key=config['comet_config']['api_key'],
project_name=config['comet_config']['project_name'],
workspace=config['comet_config']['workspace'],
)
comet_logger.add_tag(config['comet_tag'])
comet_logger.set_name(config['comet_name'])
comet_logger.log_parameters(config)
print("Comet Logger Initialized.")
comet_logger = create_comet_experiment(config)
print("Comet Logger Initialized.")

dist.barrier()

Expand Down
22 changes: 7 additions & 15 deletions finetune/train_tokenizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,23 +10,22 @@
import torch.nn.functional as F
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
from torch.nn.parallel import DistributedDataParallel as DDP

import comet_ml
from torch.nn.parallel import DistributedDataParallel as DDP

# Ensure project root is in path
sys.path.append("../")
from config import Config
from dataset import QlibDataset
from model.kronos import KronosTokenizer
# Import shared utilities
from utils.training_utils import (
from utils.training_utils import (
setup_ddp,
cleanup_ddp,
set_seed,
get_model_size,
format_time,
)
format_time,
)
from utils.experiment_logging import create_comet_experiment


def create_dataloaders(config: dict, rank: int, world_size: int):
Expand Down Expand Up @@ -235,15 +234,8 @@ def main(config: dict):
'world_size': world_size,
}
if config['use_comet']:
comet_logger = comet_ml.Experiment(
api_key=config['comet_config']['api_key'],
project_name=config['comet_config']['project_name'],
workspace=config['comet_config']['workspace'],
)
comet_logger.add_tag(config['comet_tag'])
comet_logger.set_name(config['comet_name'])
comet_logger.log_parameters(config)
print("Comet Logger Initialized.")
comet_logger = create_comet_experiment(config)
print("Comet Logger Initialized.")

dist.barrier() # Ensure save directory is created before proceeding

Expand Down
36 changes: 36 additions & 0 deletions finetune/utils/experiment_logging.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,36 @@
"""Optional experiment-tracking integrations for fine-tuning."""

import importlib
from typing import Any, Mapping, Optional


def create_comet_experiment(config: Mapping[str, Any]) -> Optional[Any]:
"""Create a Comet experiment only when Comet logging is enabled.

Keeping the import here allows fine-tuning users who do not use Comet to
run without installing its optional dependency.
"""
if not config.get("use_comet", False):
return None

try:
comet_ml = importlib.import_module("comet_ml")
except ModuleNotFoundError as exc:
if exc.name == "comet_ml":
raise RuntimeError(
"Comet logging is enabled, but the optional dependency is not "
"installed. Install it with `pip install comet_ml`, or set "
"use_comet = False` in finetune/config.py."
) from exc
raise

comet_config = config["comet_config"]
experiment = comet_ml.Experiment(
api_key=comet_config["api_key"],
project_name=comet_config["project_name"],
workspace=comet_config["workspace"],
)
experiment.add_tag(config["comet_tag"])
experiment.set_name(config["comet_name"])
experiment.log_parameters(dict(config))
return experiment
60 changes: 60 additions & 0 deletions tests/test_experiment_logging.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,60 @@
import importlib.util
from pathlib import Path
from types import SimpleNamespace
import unittest
from unittest.mock import Mock, patch


MODULE_PATH = Path(__file__).parents[1] / "finetune" / "utils" / "experiment_logging.py"
SPEC = importlib.util.spec_from_file_location("experiment_logging", MODULE_PATH)
experiment_logging = importlib.util.module_from_spec(SPEC)
SPEC.loader.exec_module(experiment_logging)


class CreateCometExperimentTests(unittest.TestCase):
def setUp(self):
self.config = {
"use_comet": True,
"comet_config": {
"api_key": "test-key",
"project_name": "test-project",
"workspace": "test-workspace",
},
"comet_tag": "test-tag",
"comet_name": "test-name",
}

def test_disabled_logging_does_not_import_comet(self):
config = {"use_comet": False}
with patch.object(experiment_logging.importlib, "import_module") as import_module:
self.assertIsNone(experiment_logging.create_comet_experiment(config))
import_module.assert_not_called()

def test_enabled_logging_explains_missing_optional_dependency(self):
with patch.object(
experiment_logging.importlib,
"import_module",
side_effect=ModuleNotFoundError("No module named 'comet_ml'", name="comet_ml"),
):
with self.assertRaisesRegex(RuntimeError, "pip install comet_ml"):
experiment_logging.create_comet_experiment(self.config)

def test_enabled_logging_configures_experiment(self):
experiment = Mock()
comet_module = SimpleNamespace(Experiment=Mock(return_value=experiment))
with patch.object(experiment_logging.importlib, "import_module", return_value=comet_module):
result = experiment_logging.create_comet_experiment(self.config)

self.assertIs(result, experiment)
comet_module.Experiment.assert_called_once_with(
api_key="test-key",
project_name="test-project",
workspace="test-workspace",
)
experiment.add_tag.assert_called_once_with("test-tag")
experiment.set_name.assert_called_once_with("test-name")
experiment.log_parameters.assert_called_once_with(self.config)


if __name__ == "__main__":
unittest.main()