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
11 changes: 9 additions & 2 deletions dspy/primitives/base_module.py
Original file line number Diff line number Diff line change
Expand Up @@ -281,11 +281,18 @@ def load(self, path, allow_pickle=False, allow_unsafe_lm_state=False):
dependency_versions = get_dependency_versions()
saved_dependency_versions = state["metadata"]["dependency_versions"]
for key, saved_version in saved_dependency_versions.items():
if dependency_versions[key] != saved_version:
current_version = dependency_versions.get(key)
if current_version is None:
logger.warning(
f"Saved model references dependency '{key}' (saved version {saved_version}) that is not tracked "
"in the current environment. This might cause errors or performance downgrade on the loaded model."
)
continue
if current_version != saved_version:
logger.warning(
f"There is a mismatch of {key} version between saved model and current environment. "
f"You saved with `{key}=={saved_version}`, but now you have "
f"`{key}=={dependency_versions[key]}`. This might cause errors or performance downgrade "
f"`{key}=={current_version}`. This might cause errors or performance downgrade "
"on the loaded model, please consider loading the model in the same environment as the "
"saving environment."
)
Expand Down
15 changes: 11 additions & 4 deletions dspy/utils/saving.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,12 +49,19 @@ def load(path: str, allow_pickle: bool = False) -> "Module":
dependency_versions = get_dependency_versions()
saved_dependency_versions = metadata["dependency_versions"]
for key, saved_version in saved_dependency_versions.items():
if dependency_versions[key] != saved_version:
current_version = dependency_versions.get(key)
if current_version is None:
logger.warning(
f"Saved model references dependency '{key}' (saved version {saved_version}) that is not tracked in the "
"current environment. This might cause errors or performance downgrade on the loaded model."
)
continue
if current_version != saved_version:
logger.warning(
f"There is a mismatch of {key} version between saved model and current environment. You saved with "
f"`{key}=={saved_version}`, but now you have `{key}=={dependency_versions[key]}`. This might cause "
"errors or performance downgrade on the loaded model, please consider loading the model in the same "
"environment as the saving environment."
f"`{key}=={saved_version}`, but now you have `{key}=={current_version}`. This might cause errors or "
"performance downgrade on the loaded model, please consider loading the model in the same environment "
"as the saving environment."
)

with open(path / "program.pkl", "rb") as f:
Expand Down
53 changes: 53 additions & 0 deletions tests/primitives/test_base_module.py
Original file line number Diff line number Diff line change
Expand Up @@ -245,6 +245,59 @@ def emit(self, record):
logger.removeHandler(handler)


def test_json_load_with_unknown_dependency_key(tmp_path):
from dspy.primitives.base_module import logger

# Saved metadata tracks an extra dependency the loading environment does not know about.
save_versions = {"python": "3.10", "dspy": "2.5.0", "cloudpickle": "2.0", "numpy": "1.26.0"}

# Loading environment tracks the standard three keys only (numpy absent).
load_versions = {"python": "3.10", "dspy": "2.5.0", "cloudpickle": "2.0"}

predict = dspy.Predict("question->answer")

class ListHandler(logging.Handler):
def __init__(self):
super().__init__()
self.messages = []

def emit(self, record):
self.messages.append(record.getMessage())

handler = ListHandler()
original_level = logger.level
logger.addHandler(handler)
logger.setLevel(logging.WARNING)

try:
save_path = tmp_path / "program.json"
# Mock version during save (with the extra numpy key)
with patch("dspy.primitives.base_module.get_dependency_versions", return_value=save_versions):
predict.save(save_path)

# Mock version during load (without the numpy key) — previously raised KeyError.
# JSON path requires no `allow_pickle=True`.
with patch("dspy.primitives.base_module.get_dependency_versions", return_value=load_versions):
loaded_predict = dspy.Predict("question->answer")
loaded_predict.load(save_path)

# Exactly one warning: the untracked numpy key. The known keys match, so no mismatch warnings.
# The JSON path logs no pickle warning, unlike the .pkl path.
assert len(handler.messages) == 1
assert "numpy" in handler.messages[0]
assert "not tracked" in handler.messages[0]

# Verify the model still loads correctly despite the unknown dependency key
assert isinstance(loaded_predict, dspy.Predict)
assert str(predict.signature) == str(loaded_predict.signature)
assert loaded_predict.dump_state() == predict.dump_state()

finally:
# Clean up: restore original level and remove handler
logger.setLevel(original_level)
logger.removeHandler(handler)


@pytest.mark.llm_call
def test_single_module_call_with_usage_tracker(lm_for_test):
dspy.configure(lm=dspy.LM(lm_for_test, cache=False, temperature=0.0), track_usage=True)
Expand Down
48 changes: 48 additions & 0 deletions tests/utils/test_saving.py
Original file line number Diff line number Diff line change
Expand Up @@ -133,6 +133,54 @@ def emit(self, record):
logger.removeHandler(handler)


def test_load_with_unknown_dependency_key(tmp_path):
from dspy.utils.saving import logger

# Saved metadata tracks an extra dependency the loading environment does not know about.
save_versions = {"python": "3.10", "dspy": "2.5.0", "cloudpickle": "2.0", "numpy": "1.26.0"}

# Loading environment tracks the standard three keys only (numpy absent).
load_versions = {"python": "3.10", "dspy": "2.5.0", "cloudpickle": "2.0"}

predict = dspy.Predict("question->answer")

class ListHandler(logging.Handler):
def __init__(self):
super().__init__()
self.messages = []

def emit(self, record):
self.messages.append(record.getMessage())

handler = ListHandler()
original_level = logger.level
logger.addHandler(handler)
logger.setLevel(logging.WARNING)

try:
# Mock version during save (with the extra numpy key)
with patch("dspy.primitives.base_module.get_dependency_versions", return_value=save_versions):
predict.save(tmp_path, save_program=True)

# Mock version during load (without the numpy key) — previously raised KeyError
with patch("dspy.utils.saving.get_dependency_versions", return_value=load_versions):
loaded_predict = dspy.load(tmp_path, allow_pickle=True)

# Exactly one warning: the untracked numpy key. The known keys match, so no mismatch warnings.
assert len(handler.messages) == 1
assert "numpy" in handler.messages[0]
assert "not tracked" in handler.messages[0]

# Verify the model still loads correctly despite the unknown dependency key
assert isinstance(loaded_predict, dspy.Predict)
assert predict.signature == loaded_predict.signature

finally:
# Clean up: restore original level and remove handler
logger.setLevel(original_level)
logger.removeHandler(handler)


def test_pickle_loading_requires_explicit_permission(tmp_path):
"""Test that loading pickle files requires explicit permission."""
predict = dspy.Predict("question->answer")
Expand Down