Skip to content
Merged
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
6 changes: 6 additions & 0 deletions .gitattributes
Original file line number Diff line number Diff line change
@@ -0,0 +1,6 @@
# The sources the ITK style checks read (KWStyle, clang-format), as ITK marks its own.
[attr]our-c-style whitespace=tab-in-indent,no-lf-at-eof hooks.style=KWStyle,clangformat

*.h our-c-style
*.cxx our-c-style
*.hxx our-c-style
18 changes: 9 additions & 9 deletions examples/ImpactMetricExample.cxx
Original file line number Diff line number Diff line change
Expand Up @@ -49,11 +49,11 @@ main(int argc, char * argv[])
std::cerr << " device: \"cpu\" (default), \"cuda\", \"cuda:0\", ...\n";
return EXIT_FAILURE;
}
const char * const modelPath = argv[1];
const char * const fixedPath = argv[2];
const char * const movingPath = argv[3];
const char * const outputPath = argv[4];
const std::string device = (argc > 5) ? argv[5] : "cpu";
const char * const modelPath = argv[1];
const char * const fixedPath = argv[2];
const char * const movingPath = argv[3];
const char * const outputPath = argv[4];
const std::string device = (argc > 5) ? argv[5] : "cpu";

constexpr unsigned int Dimension = 3;
using PixelType = float;
Expand Down Expand Up @@ -84,14 +84,14 @@ main(int argc, char * argv[])

// --- 2. metric: register moving onto fixed by comparing anatomical features -------
using MetricType = itk::ImpactImageToImageMetricv4<ImageType, ImageType>;
auto metric = MetricType::New();
auto metric = MetricType::New();
std::vector<itk::ImpactModelConfiguration> models{ config };
metric->SetModelsConfiguration(models);
metric->SetDistance({ "L2" }); // per-layer loss: L1, L2, NCC, Cosine, Dice, ...
metric->SetDistance({ "L2" }); // per-layer loss: L1, L2, NCC, Cosine, Dice, ...
metric->SetLayersWeight({ 1.0f });
metric->SetSubsetFeatures({ 4 }); // random channel subset for speed (0 = all)
metric->SetSubsetFeatures({ 4 }); // random channel subset for speed (0 = all)
metric->SetPCA({ 0 });
metric->SetMode("Static"); // "Static" (precomputed features) or "Jacobian"
metric->SetMode("Static"); // "Static" (precomputed features) or "Jacobian"
metric->SetDevice(device);

using TransformType = itk::TranslationTransform<double, Dimension>;
Expand Down
6 changes: 2 additions & 4 deletions include/ImpactLoss.h
Original file line number Diff line number Diff line change
Expand Up @@ -537,8 +537,7 @@ class L1Cosine : public Loss
}
};

inline RegisterLoss<L1Cosine> L1Cosine_reg(
"L1Cosine");
inline RegisterLoss<L1Cosine> L1Cosine_reg("L1Cosine");

/**
* \class Cosine
Expand Down Expand Up @@ -763,8 +762,7 @@ class NCC : public Loss
if (N <= 0)
return 0.0;
torch::Tensor u = this->m_sfm - (this->m_sf * this->m_sm / N);
torch::Tensor v =
Denominator(Variance(this->m_sff, this->m_sf, N), Variance(this->m_smm, this->m_sm, N));
torch::Tensor v = Denominator(Variance(this->m_sff, this->m_sf, N), Variance(this->m_smm, this->m_sm, N));
return 1.0 - (u / v).mean().item<double>();
}

Expand Down
5 changes: 2 additions & 3 deletions include/itkImageToFeaturesMap.hxx
Original file line number Diff line number Diff line change
Expand Up @@ -53,14 +53,13 @@ ImageToFeaturesMap<TInputImage, TInterpolator>::ImageToFeaturesMap()

template <typename TInputImage, typename TInterpolator>
void
ImageToFeaturesMap<TInputImage, TInterpolator>
::PrintSelf(std::ostream & os, Indent indent) const
ImageToFeaturesMap<TInputImage, TInterpolator>::PrintSelf(std::ostream & os, Indent indent) const
{
Superclass::PrintSelf(os, indent);
}

template <typename TInputImage, typename TInterpolator>
void
void
ImageToFeaturesMap<TInputImage, TInterpolator>::AddInput(const TInputImage * input)
{
if (!m_Interpolator)
Expand Down
24 changes: 15 additions & 9 deletions include/itkImageToTensorFilter.h
Original file line number Diff line number Diff line change
Expand Up @@ -54,10 +54,10 @@ class ITK_TEMPLATE_EXPORT ImageToTensorFilter : public ProcessObject
using InputImagePointer = typename InputImageType::Pointer;
using InputImageConstPointer = typename InputImageType::ConstPointer;

using TensorType = torch::Tensor;
using TensorType = torch::Tensor;
using TensorHandleType = std::shared_ptr<TensorType>;
using TensorDataObject = itk::SimpleDataObjectDecorator<TensorHandleType>;

using InterpolatorPointer = typename TInterpolator::Pointer;

using InputImagePixelType = typename InputImageType::PixelType;
Expand Down Expand Up @@ -89,9 +89,14 @@ class ITK_TEMPLATE_EXPORT ImageToTensorFilter : public ProcessObject
itkSetMacro(OutputSpacing, InputSpacingType);
itkGetConstReferenceMacro(OutputSpacing, InputSpacingType);
itkSetVectorMacro(OutputSpacing, const float, ImageDimension);

void SetInterpolator(typename TInterpolator::Pointer interp) { m_Interpolator = interp; }
void SetTransform(std::function<InputImagePointType(const InputImagePointType &)> fct)

void
SetInterpolator(typename TInterpolator::Pointer interp)
{
m_Interpolator = interp;
}
void
SetTransform(std::function<InputImagePointType(const InputImagePointType &)> fct)
{
m_Transform = fct;
}
Expand Down Expand Up @@ -141,15 +146,16 @@ class ITK_TEMPLATE_EXPORT ImageToTensorFilter : public ProcessObject

void
PrintSelf(std::ostream & os, Indent indent) const override;

void
VerifyPreconditions() const override;

void GenerateData() override;
void
GenerateData() override;

private:
InterpolatorPointer m_Interpolator;
InputSpacingType m_OutputSpacing{ MakeFilled<InputSpacingType>(1.0) };
InterpolatorPointer m_Interpolator;
InputSpacingType m_OutputSpacing{ MakeFilled<InputSpacingType>(1.0) };
std::function<InputImagePointType(const InputImagePointType &)> m_Transform;
};

Expand Down
14 changes: 7 additions & 7 deletions include/itkImpactCoarseRegistration.hxx
Original file line number Diff line number Diff line change
Expand Up @@ -276,8 +276,7 @@ ImpactCoarseRegistration<TFixedImage, TMovingImage>::GetDisplacementField() -> D

template <typename TFixedImage, typename TMovingImage>
auto
ImpactCoarseRegistration<TFixedImage, TMovingImage>::GetDisplacementFieldTransform()
-> DisplacementFieldTransformType *
ImpactCoarseRegistration<TFixedImage, TMovingImage>::GetDisplacementFieldTransform() -> DisplacementFieldTransformType *
{
return m_DisplacementFieldTransform.GetPointer();
}
Expand Down Expand Up @@ -421,7 +420,9 @@ ImpactCoarseRegistration<TFixedImage, TMovingImage>::GenerateData()
if (kept > 0 && kept < fixedLayers[l].size(1))
{
const torch::Tensor channels =
torch::randperm(fixedLayers[l].size(1), torch::TensorOptions().dtype(torch::kLong)).narrow(0, 0, kept).to(device);
torch::randperm(fixedLayers[l].size(1), torch::TensorOptions().dtype(torch::kLong))
.narrow(0, 0, kept)
.to(device);
fixedLayers[l] = fixedLayers[l].index_select(1, channels).contiguous();
movingLayers[l] = movingLayers[l].index_select(1, channels).contiguous();
}
Expand Down Expand Up @@ -741,8 +742,7 @@ ImpactCoarseRegistration<TFixedImage, TMovingImage>::GenerateData()
torch::Tensor f2 = (dispBack / scaleT).flip(1);

// Identity sampling grid at coarse resolution, channel-first {1, Dim, coarse...}, x,y,z.
torch::Tensor idAffine =
torch::eye(ImageDimension, torch::TensorOptions().dtype(torch::kFloat32).device(device));
torch::Tensor idAffine = torch::eye(ImageDimension, torch::TensorOptions().dtype(torch::kFloat32).device(device));
idAffine =
torch::cat({ idAffine, torch::zeros({ static_cast<int64_t>(ImageDimension), 1 }, idAffine.options()) }, 1)
.unsqueeze(0);
Expand All @@ -753,7 +753,7 @@ ImpactCoarseRegistration<TFixedImage, TMovingImage>::GenerateData()
{
gridSize.push_back(c);
}
torch::Tensor idGrid = torch::affine_grid_generator(idAffine, gridSize, /*align_corners=*/true);
torch::Tensor idGrid = torch::affine_grid_generator(idAffine, gridSize, /*align_corners=*/true);
std::vector<int64_t> toChannelFirst;
toChannelFirst.push_back(0);
toChannelFirst.push_back(idGrid.dim() - 1);
Expand Down Expand Up @@ -803,7 +803,7 @@ ImpactCoarseRegistration<TFixedImage, TMovingImage>::GenerateData()
std::vector<int64_t> componentShape(ImageDimension + 2, 1); // {1, Dim, 1, ...}: one factor per component
componentShape[1] = static_cast<int64_t>(ImageDimension);
const torch::Tensor cellsToVoxels = torch::tensor(cellVoxels, torch::kLong).to(disp.options()).view(componentShape);
torch::Tensor dispFull;
torch::Tensor dispFull;
if constexpr (ImageDimension == 3)
dispFull = F::interpolate(
disp * cellsToVoxels, F::InterpolateFuncOptions().size(exactSize).mode(torch::kTrilinear).align_corners(false));
Expand Down
18 changes: 9 additions & 9 deletions include/itkImpactFineRegistration.h
Original file line number Diff line number Diff line change
Expand Up @@ -287,18 +287,18 @@ class ITK_TEMPLATE_EXPORT ImpactFineRegistration
GenerateData() override;

private:
typename FixedImageType::ConstPointer m_FixedImage{ nullptr };
typename MovingImageType::ConstPointer m_MovingImage{ nullptr };
typename MaskImageType::ConstPointer m_FixedMask{ nullptr };
typename MaskImageType::ConstPointer m_MovingMask{ nullptr };
typename DisplacementFieldType::Pointer m_InitialDisplacementField{ nullptr };
typename FixedImageType::ConstPointer m_FixedImage{ nullptr };
typename MovingImageType::ConstPointer m_MovingImage{ nullptr };
typename MaskImageType::ConstPointer m_FixedMask{ nullptr };
typename MaskImageType::ConstPointer m_MovingMask{ nullptr };
typename DisplacementFieldType::Pointer m_InitialDisplacementField{ nullptr };

std::vector<ImpactModelConfiguration> m_FixedModelsConfiguration;
std::vector<ImpactModelConfiguration> m_MovingModelsConfiguration;
std::vector<std::string> m_Distance;
std::vector<float> m_LayersWeight;
std::vector<unsigned int> m_SubsetFeatures;
std::vector<unsigned int> m_PCA;
std::vector<std::string> m_Distance;
std::vector<float> m_LayersWeight;
std::vector<unsigned int> m_SubsetFeatures;
std::vector<unsigned int> m_PCA;

std::string m_Device{ "cpu" };
unsigned int m_Seed{ 0 };
Expand Down
67 changes: 34 additions & 33 deletions include/itkImpactFineRegistration.hxx
Original file line number Diff line number Diff line change
Expand Up @@ -181,8 +181,7 @@ ImpactFineRegistration<TFixedImage, TMovingImage>::GetDisplacementField() -> Dis

template <typename TFixedImage, typename TMovingImage>
auto
ImpactFineRegistration<TFixedImage, TMovingImage>::GetDisplacementFieldTransform()
-> DisplacementFieldTransformType *
ImpactFineRegistration<TFixedImage, TMovingImage>::GetDisplacementFieldTransform() -> DisplacementFieldTransformType *
{
return m_DisplacementFieldTransform.GetPointer();
}
Expand Down Expand Up @@ -311,25 +310,27 @@ ImpactFineRegistration<TFixedImage, TMovingImage>::GenerateData()
if (coarseSpatial != spatial)
{
if constexpr (ImageDimension == 3)
initField = torch::nn::functional::interpolate(
initField,
torch::nn::functional::InterpolateFuncOptions().size(coarseSpatial).mode(torch::kTrilinear).align_corners(true));
initField = torch::nn::functional::interpolate(initField,
torch::nn::functional::InterpolateFuncOptions()
.size(coarseSpatial)
.mode(torch::kTrilinear)
.align_corners(true));
else
initField = torch::nn::functional::interpolate(
initField,
torch::nn::functional::InterpolateFuncOptions().size(coarseSpatial).mode(torch::kBilinear).align_corners(true));
initField = torch::nn::functional::interpolate(initField,
torch::nn::functional::InterpolateFuncOptions()
.size(coarseSpatial)
.mode(torch::kBilinear)
.align_corners(true));
}
theta = initField.set_requires_grad(true);
}
else
{
theta =
torch::zeros(fieldShape, torch::TensorOptions().dtype(torch::kFloat32).device(device).requires_grad(true));
theta = torch::zeros(fieldShape, torch::TensorOptions().dtype(torch::kFloat32).device(device).requires_grad(true));
}

// ---- 3. Base identity sampling grid (normalized [-1,1], last-dim order x,y,z), align_corners=true. ----
torch::Tensor idAffine =
torch::eye(ImageDimension, torch::TensorOptions().dtype(torch::kFloat32).device(device));
torch::Tensor idAffine = torch::eye(ImageDimension, torch::TensorOptions().dtype(torch::kFloat32).device(device));
idAffine = torch::cat({ idAffine, torch::zeros({ static_cast<int64_t>(ImageDimension), 1 }, idAffine.options()) }, 1)
.unsqueeze(0); // {1, N, N+1}
std::vector<int64_t> gridSize;
Expand All @@ -339,8 +340,7 @@ ImpactFineRegistration<TFixedImage, TMovingImage>::GenerateData()
{
gridSize.push_back(s);
}
torch::Tensor grid0 =
torch::affine_grid_generator(idAffine, gridSize, /*align_corners=*/true); // {1, z,y,x, N}
torch::Tensor grid0 = torch::affine_grid_generator(idAffine, gridSize, /*align_corners=*/true); // {1, z,y,x, N}

// Per-component normalization in (z, y, x) order: the field units (s_min) per grid_sample unit, (size-1)/2 voxels
// (exact for align_corners=true) of s_a / s_min units each.
Expand Down Expand Up @@ -394,7 +394,8 @@ ImpactFineRegistration<TFixedImage, TMovingImage>::GenerateData()
return field;
}
if constexpr (ImageDimension == 3)
return F::interpolate(field, F::InterpolateFuncOptions().size(target).mode(torch::kTrilinear).align_corners(true));
return F::interpolate(field,
F::InterpolateFuncOptions().size(target).mode(torch::kTrilinear).align_corners(true));
else
return F::interpolate(field, F::InterpolateFuncOptions().size(target).mode(torch::kBilinear).align_corners(true));
};
Expand Down Expand Up @@ -510,15 +511,15 @@ ImpactFineRegistration<TFixedImage, TMovingImage>::GenerateData()
// Feature-mode setup: extract the fixed/moving feature layers (constants, not differentiated
// through), optionally PCA-reduce them (fit on fixed), and build one loss per kept layer.
// Intensity mode skips this and compares raw voxels.
std::vector<torch::Tensor> fixedLayers;
std::vector<torch::Tensor> movingLayers;
std::vector<torch::Tensor> fixedLayersOnline; // "Jacobian" mode: fixed features, pooled
std::vector<torch::Tensor> pcaBasis; // per kept layer; undefined entry = no PCA
std::vector<std::unique_ptr<Impact::Loss>> losses;
std::vector<float> layerWeights;
std::vector<torch::Tensor> fixedLayers;
std::vector<torch::Tensor> movingLayers;
std::vector<torch::Tensor> fixedLayersOnline; // "Jacobian" mode: fixed features, pooled
std::vector<torch::Tensor> pcaBasis; // per kept layer; undefined entry = no PCA
std::vector<std::unique_ptr<Impact::Loss>> losses;
std::vector<float> layerWeights;
// SubsetFeatures: per kept layer, that many of its channels drawn at random at every iteration (0 = all).
std::vector<torch::Tensor> subsets; // per layer; an undefined entry keeps every channel
auto drawSubsets = [&](const std::vector<torch::Tensor> & layers) {
auto drawSubsets = [&](const std::vector<torch::Tensor> & layers) {
subsets.assign(layers.size(), torch::Tensor());
for (size_t l = 0; l < layers.size(); ++l)
{
Expand Down Expand Up @@ -731,8 +732,7 @@ ImpactFineRegistration<TFixedImage, TMovingImage>::GenerateData()
}
for (size_t l = 0; l < layerCount; ++l)
{
const std::string name =
m_Distance.empty() ? std::string("L2") : m_Distance[std::min(l, m_Distance.size() - 1)];
const std::string name = m_Distance.empty() ? std::string("L2") : m_Distance[std::min(l, m_Distance.size() - 1)];
losses.push_back(Impact::LossFactory::Instance().Create(name));
layerWeights.push_back(l < m_LayersWeight.size() ? m_LayersWeight[l] : 1.0f);
if (sampled && losses.back()->IsSpatial())
Expand Down Expand Up @@ -951,18 +951,16 @@ ImpactFineRegistration<TFixedImage, TMovingImage>::GenerateData()
}

// ---- 4. Adam loop (entirely on device; no host copies). ----
torch::optim::Adam optimizer({ theta },
torch::optim::AdamOptions(m_LearningRate)
.betas(std::make_tuple(m_Beta1, m_Beta2))
.eps(m_Epsilon));
torch::optim::Adam optimizer(
{ theta }, torch::optim::AdamOptions(m_LearningRate).betas(std::make_tuple(m_Beta1, m_Beta2)).eps(m_Epsilon));

m_MetricValuesPerIteration.clear();
m_MetricValuesPerIteration.reserve(m_NumberOfIterations);

// Every layer's loss divided by its value at the stage's first iteration, so each starts at 1 and
// LayersWeight weighs comparable quantities (see Impact::LossNormalization).
Impact::LossNormalization normalization;
auto normalized = [&](size_t l, const torch::Tensor & value) -> torch::Tensor {
auto normalized = [&](size_t l, const torch::Tensor & value) -> torch::Tensor {
if (!m_NormalizeLosses)
{
return value;
Expand Down Expand Up @@ -1032,7 +1030,7 @@ ImpactFineRegistration<TFixedImage, TMovingImage>::GenerateData()
const torch::Tensor smoothedControl = smoothControl(theta); // graph to theta
torch::Tensor smoothed = smoothedControl.detach().requires_grad_(true);
smoothed.mutable_grad() = torch::zeros_like(smoothed);
int64_t voxels = 1;
int64_t voxels = 1;
for (const int64_t extent : spatial)
{
voxels *= extent;
Expand Down Expand Up @@ -1439,9 +1437,11 @@ ImpactFineRegistration<TFixedImage, TMovingImage>::GenerateData()
if (cur != tgt)
{
if constexpr (ImageDimension == 3)
mll = F::interpolate(mll, F::InterpolateFuncOptions().size(tgt).mode(torch::kTrilinear).align_corners(true));
mll =
F::interpolate(mll, F::InterpolateFuncOptions().size(tgt).mode(torch::kTrilinear).align_corners(true));
else
mll = F::interpolate(mll, F::InterpolateFuncOptions().size(tgt).mode(torch::kBilinear).align_corners(true));
mll =
F::interpolate(mll, F::InterpolateFuncOptions().size(tgt).mode(torch::kBilinear).align_corners(true));
}
torch::Tensor counted;
if (masked)
Expand Down Expand Up @@ -1575,7 +1575,8 @@ ImpactFineRegistration<TFixedImage, TMovingImage>::GenerateData()
m_WarpedMovingImage->SetDirection(m_FixedImage->GetDirection());
m_WarpedMovingImage->Allocate();

ImageRegionIteratorWithIndex<WarpedImageType> wit(m_WarpedMovingImage, m_WarpedMovingImage->GetLargestPossibleRegion());
ImageRegionIteratorWithIndex<WarpedImageType> wit(m_WarpedMovingImage,
m_WarpedMovingImage->GetLargestPossibleRegion());
for (wit.GoToBegin(); !wit.IsAtEnd(); ++wit)
{
const auto idx = wit.GetIndex();
Expand Down
14 changes: 7 additions & 7 deletions include/itkImpactImageToImageMetricv4.h
Original file line number Diff line number Diff line change
Expand Up @@ -303,18 +303,18 @@ class ITK_TEMPLATE_EXPORT ImpactImageToImageMetricv4
int m_FeaturesMapUpdateInterval{ 0 };
/** Derivative evaluations since Initialize(), the metric's stand-in for an iteration count.
* Mutable because GetValueAndDerivative() is const, as the base metric declares it. */
mutable unsigned long m_CurrentIteration{ 0 };
std::string m_Mode;
std::string m_FeatureMapsPath;
bool m_NormalizeLosses{ true };
mutable unsigned long m_CurrentIteration{ 0 };
std::string m_Mode;
std::string m_FeatureMapsPath;
bool m_NormalizeLosses{ true };
/** Latched at the first evaluation after Initialize(); mutable because the evaluation is const. */
mutable Impact::LossNormalization m_LossNormalization;
std::string m_Device = "cpu";
std::string m_Device = "cpu";
// Zero means "seed from the clock", which is what the per-work-unit generator does with it.
// It must have a value even when the user never calls SetSeed(): it is read on every
// evaluation, and an indeterminate one would make the metric unreproducible at random.
unsigned int m_Seed{ 0 };
unsigned int m_BatchSize{ 64 };
unsigned int m_Seed{ 0 };
unsigned int m_BatchSize{ 64 };

std::vector<std::vector<unsigned int>> m_features_indexes;
};
Expand Down
Loading
Loading