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
25 changes: 25 additions & 0 deletions ptf/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,25 @@
from importlib.abc import MetaPathFinder
import sys


# Custom import interceptor to redirect all sub-imports
# (e.g., ptf.models -> pytorch_forecasting.models)
class LegacyRedirectFinder(MetaPathFinder):
def find_spec(self, fullname, path, target=None):
if fullname.startswith("ptf"):
real_name = fullname.replace("ptf", "pytorch_forecasting", 1)
try:
__import__(real_name)
sys.modules[fullname] = sys.modules[real_name]
return sys.modules[real_name].__spec__
except ImportError:
return None
return None


sys.meta_path.insert(0, LegacyRedirectFinder())

# Map the root package to pytorch_forecasting
import pytorch_forecasting

sys.modules["ptf"] = pytorch_forecasting
16 changes: 16 additions & 0 deletions tests/test_ptf_import.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,16 @@
import sys


def test_ptf_import_redirection():
# Clear any previous ptf imports from sys.modules
for module_name in list(sys.modules.keys()):
if module_name.startswith("ptf"):
del sys.modules[module_name]

# Import from ptf
import ptf.models as ptf_models
import pytorch_forecasting.models as pf_models

# Verify redirection and content equality
assert hasattr(ptf_models, "TemporalFusionTransformer")
assert ptf_models.TemporalFusionTransformer is pf_models.TemporalFusionTransformer
Loading