Skip to content
Draft
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
215 changes: 108 additions & 107 deletions ctlearn/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,17 +7,20 @@
import numpy as np
import pytest
import shutil
from unittest import mock
from astropy import units as u
from astropy.table import Column, Table
from traitlets.config.loader import Config

from ctapipe.core import run_tool
from ctapipe.io import write_table
from ctapipe.utils import get_dataset_path
from ctlearn.tools import DLFrameWork
from ctlearn.tools.keras.train_model import TrainCTLearnKerasModel
from ctlearn.utils import get_lst1_subarray_description

# TODO: ADD PyTorch here
TRAINING_TOOLS = {"Keras": TrainCTLearnKerasModel}
MODEL_FILE_FORMATS = {"Keras": "keras", "PyTorch": "pth"}


@pytest.fixture(scope="session")
def gamma_simtel_path():
Expand Down Expand Up @@ -241,41 +244,41 @@ def ctlearn_trained_r1_mono_models(r1_gamma_file, r1_proton_file, tmp_path_facto
telescope_type = "LST"
# Loop over reconstruction tasks and train models for each combination
ctlearn_trained_r1_mono_models = {}
with mock.patch("ctapipe.instrument.SubarrayDescription.__eq__", return_value=True):
for reco_task in ["type", "energy", "cameradirection"]:
# Output directory for trained model
output_dir = tmp_path / f"ctlearn_{telescope_type}_{reco_task}"

# Build command-line arguments
argv = [
f"--signal={signal_dir}",
"--pattern-signal=*.r1.h5",
f"--output={output_dir}",
f"--reco={reco_task}",
"--TrainCTLearnModel.n_epochs=1",
"--TrainCTLearnModel.batch_size=2",
"--TrainCTLearnModel.dl1dh_reader_type=DLWaveformReader",
"--DLWaveformReader.sequence_length=5",
"--DLWaveformReader.focal_length_choice=EQUIVALENT",
]

# Include background only for classification task
if reco_task == "type":
argv.extend(
[
f"--background={background_dir}",
"--pattern-background=*.r1.h5",
]
)

# Run training
assert run_tool(DLFrameWork(config=config), argv=argv, cwd=tmp_path) == 0
for reco_task in ["type", "energy", "cameradirection"]:
# Output directory for trained model
output_dir = tmp_path / f"ctlearn_{telescope_type}_{reco_task}"

# Build command-line arguments
argv = [
f"--signal={signal_dir}",
"--pattern-signal=*.r1.h5",
f"--output={output_dir}",
f"--reco={reco_task}",
"--TrainCTLearnModel.n_epochs=1",
"--TrainCTLearnModel.batch_size=2",
"--TrainCTLearnModel.dl1dh_reader_type=DLWaveformReader",
"--DLWaveformReader.sequence_length=5",
"--DLWaveformReader.focal_length_choice=EQUIVALENT",
]

# Include background only for classification task
if reco_task == "type":
argv.extend(
[
f"--background={background_dir}",
"--pattern-background=*.r1.h5",
"--DLWaveformReader.enforce_subarray_equality=False",
]
)

ctlearn_trained_r1_mono_models[f"{telescope_type}_{reco_task}"] = (
output_dir / "ctlearn_model.keras"
# Run training tools
for framework, training_tool in TRAINING_TOOLS.items():
assert run_tool(training_tool(config=config), argv=argv, cwd=tmp_path) == 0
ctlearn_trained_r1_mono_models[f"{framework}_{telescope_type}_{reco_task}"] = (
output_dir / f"ctlearn_model.{MODEL_FILE_FORMATS[framework]}"
)
# Check that the trained model exists
assert ctlearn_trained_r1_mono_models[f"{telescope_type}_{reco_task}"].exists()
assert ctlearn_trained_r1_mono_models[f"{framework}_{telescope_type}_{reco_task}"].exists()
return ctlearn_trained_r1_mono_models


Expand Down Expand Up @@ -323,47 +326,45 @@ def ctlearn_trained_dl1_mono_models(dl1_gamma_file, dl1_proton_file, tmp_path_fa
# Loop over telescope types and reconstruction tasks
# and train models for each combination
ctlearn_trained_dl1_mono_models = {}
with mock.patch("ctapipe.instrument.SubarrayDescription.__eq__", return_value=True):
for telescope_type, allowed_tels in telescope_types.items():
for reco_task in ["type", "energy", "cameradirection"]:
# Output directory for trained model
output_dir = tmp_path / f"ctlearn_{telescope_type}_{reco_task}"

# Build command-line arguments
argv = [
f"--signal={signal_dir}",
"--pattern-signal=*.dl1.h5",
f"--output={output_dir}",
f"--reco={reco_task}",
"--TrainCTLearnModel.n_epochs=1",
"--TrainCTLearnModel.batch_size=2",
"--DLImageReader.focal_length_choice=EQUIVALENT",
f"--DLImageReader.allowed_tels={allowed_tels}",
]

# Include background only for classification task
if reco_task == "type":
argv.extend(
[
f"--background={background_dir}",
"--pattern-background=*.dl1.h5",
f"--DLImageReader.image_mapper_type={image_mapper_types[telescope_type]}",
]
)

# Run training
assert (
run_tool(DLFrameWork(config=config), argv=argv, cwd=tmp_path) == 0
)
for telescope_type, allowed_tels in telescope_types.items():
for reco_task in ["type", "energy", "cameradirection"]:
# Output directory for trained model
output_dir = tmp_path / f"ctlearn_{telescope_type}_{reco_task}"

# Build command-line arguments
argv = [
f"--signal={signal_dir}",
"--pattern-signal=*.dl1.h5",
f"--output={output_dir}",
f"--reco={reco_task}",
"--TrainCTLearnModel.n_epochs=1",
"--TrainCTLearnModel.batch_size=2",
"--DLImageReader.focal_length_choice=EQUIVALENT",
f"--DLImageReader.allowed_tels={allowed_tels}",
]

ctlearn_trained_dl1_mono_models[f"{telescope_type}_{reco_task}"] = (
output_dir / "ctlearn_model.keras"
# Include background only for classification task
if reco_task == "type":
argv.extend(
[
f"--background={background_dir}",
"--pattern-background=*.dl1.h5",
"--DLImageReader.enforce_subarray_equality=False",
f"--DLImageReader.image_mapper_type={image_mapper_types[telescope_type]}",
]
)
# Check that the trained model exists
assert ctlearn_trained_dl1_mono_models[
f"{telescope_type}_{reco_task}"
].exists()
return ctlearn_trained_dl1_mono_models

# Run training tools
for framework, training_tool in TRAINING_TOOLS.items():
assert run_tool(training_tool(config=config), argv=argv, cwd=tmp_path) == 0
ctlearn_trained_dl1_mono_models[f"{framework}_{telescope_type}_{reco_task}"] = (
output_dir / f"ctlearn_model.{MODEL_FILE_FORMATS[framework]}"
)
# Check that the trained model exists
assert ctlearn_trained_dl1_mono_models[
f"{framework}_{telescope_type}_{reco_task}"
].exists()
return ctlearn_trained_dl1_mono_models


@pytest.fixture(scope="session")
Expand Down Expand Up @@ -404,42 +405,42 @@ def ctlearn_trained_dl1_stereo_models(

# Loop over reconstruction tasks and train models for each combination
ctlearn_trained_dl1_stereo_models = {}
with mock.patch("ctapipe.instrument.SubarrayDescription.__eq__", return_value=True):
for reco_task in ["type", "energy", "skydirection"]:
# Output directory for trained model
output_dir = tmp_path / f"ctlearn_{telescope_type}_{reco_task}"

# Build command-line arguments
argv = [
f"--signal={signal_dir}",
"--pattern-signal=*.dl1.h5",
f"--output={output_dir}",
f"--reco={reco_task}",
"--TrainCTLearnModel.n_epochs=1",
"--TrainCTLearnModel.batch_size=2",
"--TrainCTLearnModel.stack_telescope_images=True",
"--DLImageReader.mode=stereo",
"--DLImageReader.focal_length_choice=EQUIVALENT",
f"--DLImageReader.allowed_tels={allowed_tels}",
]

# Include background only for classification task
if reco_task == "type":
argv.extend(
[
f"--background={background_dir}",
"--pattern-background=*.dl1.h5",
]
)

# Run training
assert run_tool(DLFrameWork(config=config), argv=argv, cwd=tmp_path) == 0
for reco_task in ["type", "energy", "skydirection"]:
# Output directory for trained model
output_dir = tmp_path / f"ctlearn_{telescope_type}_{reco_task}"

# Build command-line arguments
argv = [
f"--signal={signal_dir}",
"--pattern-signal=*.dl1.h5",
f"--output={output_dir}",
f"--reco={reco_task}",
"--TrainCTLearnModel.n_epochs=1",
"--TrainCTLearnModel.batch_size=2",
"--TrainCTLearnModel.stack_telescope_images=True",
"--DLImageReader.mode=stereo",
"--DLImageReader.focal_length_choice=EQUIVALENT",
f"--DLImageReader.allowed_tels={allowed_tels}",
]

# Include background only for classification task
if reco_task == "type":
argv.extend(
[
f"--background={background_dir}",
"--pattern-background=*.dl1.h5",
"--DLImageReader.enforce_subarray_equality=False",
]
)

ctlearn_trained_dl1_stereo_models[f"{telescope_type}_{reco_task}"] = (
output_dir / "ctlearn_model.keras"
# Run training tools
for framework, training_tool in TRAINING_TOOLS.items():
assert run_tool(training_tool(config=config), argv=argv, cwd=tmp_path) == 0
ctlearn_trained_dl1_stereo_models[f"{framework}_{telescope_type}_{reco_task}"] = (
output_dir / f"ctlearn_model.{MODEL_FILE_FORMATS[framework]}"
)
# Check that the trained model exists
assert ctlearn_trained_dl1_stereo_models[
f"{telescope_type}_{reco_task}"
f"{framework}_{telescope_type}_{reco_task}"
].exists()
return ctlearn_trained_dl1_stereo_models
return ctlearn_trained_dl1_stereo_models
6 changes: 5 additions & 1 deletion ctlearn/core/__init__.py
Original file line number Diff line number Diff line change
@@ -1,2 +1,6 @@
"""ctlearn command line tools.
"""
ctlearn core functionalities
"""

import ctlearn.core.keras.model # noqa: F401
#import ctlearn.core.pytorch.model # noqa: F401
Loading
Loading