Skip to content

Commit 010f426

Browse files
author
Goran Jelic-Cizmek
committed
Add explicit METHOD to SOLVE if solving DERIVATIVE
1 parent a7f7753 commit 010f426

6 files changed

Lines changed: 174 additions & 0 deletions

File tree

src/nmodl/main.cpp

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -30,6 +30,7 @@
3030
#include "visitors/global_var_visitor.hpp"
3131
#include "visitors/initial_block_visitor.hpp"
3232
#include "visitors/implicit_argument_visitor.hpp"
33+
#include "visitors/implicit_method_visitor.hpp"
3334
#include "visitors/indexedname_visitor.hpp"
3435
#include "visitors/inline_visitor.hpp"
3536
#include "visitors/json_visitor.hpp"
@@ -453,6 +454,12 @@ int run_nmodl(int argc, const char* argv[]) {
453454
SymtabVisitor(update_symtab).visit_program(*ast);
454455
}
455456

457+
/// insert an explicit method to SOLVE blocks (if required)
458+
{
459+
logger->info("Running implicit method for SOLVE block visitor");
460+
ImplicitMethodVisitor().visit_program(*ast);
461+
ast_to_nmodl(*ast, filepath("implicit_solve_method"));
462+
}
456463

457464
/// note that we can not symtab visitor in update mode as we
458465
/// replace kinetic block with derivative block of same name

src/nmodl/visitors/CMakeLists.txt

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,7 @@ add_library(
1616
function_callpath_visitor.cpp
1717
global_var_visitor.cpp
1818
implicit_argument_visitor.cpp
19+
implicit_method_visitor.cpp
1920
indexedname_visitor.cpp
2021
index_remover.cpp
2122
inline_visitor.cpp
Lines changed: 43 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,43 @@
1+
/*
2+
* Copyright 2025 EPFL.
3+
* See the top-level LICENSE file for details.
4+
*
5+
* SPDX-License-Identifier: Apache-2.0
6+
*/
7+
8+
#include "ast/ast_decl.hpp"
9+
#include "ast/derivative_block.hpp"
10+
#include "ast/name.hpp"
11+
#include "ast/program.hpp"
12+
#include "ast/solve_block.hpp"
13+
#include "ast/string.hpp"
14+
#include "codegen/codegen_naming.hpp"
15+
#include "visitors/implicit_method_visitor.hpp"
16+
#include "visitors/visitor_utils.hpp"
17+
18+
namespace nmodl {
19+
namespace visitor {
20+
21+
void ImplicitMethodVisitor::visit_program(ast::Program& node) {
22+
const auto& derivative_blocks = collect_nodes(node, {ast::AstNodeType::DERIVATIVE_BLOCK});
23+
for (const auto& derivative_block: derivative_blocks) {
24+
const auto& block = std::dynamic_pointer_cast<ast::DerivativeBlock>(derivative_block);
25+
derivative_block_names.insert(block->get_node_name());
26+
}
27+
node.visit_children(*this);
28+
}
29+
30+
void ImplicitMethodVisitor::visit_solve_block(ast::SolveBlock& node) {
31+
const auto& name = node.get_block_name()->get_node_name();
32+
const auto& method = node.get_method();
33+
for (const auto& derivative_block_name: derivative_block_names) {
34+
if (derivative_block_name == name && !method) {
35+
node.set_method(std::make_shared<ast::Name>(
36+
std::make_shared<ast::String>((codegen::naming::DERIVIMPLICIT_METHOD))));
37+
}
38+
}
39+
node.visit_children(*this);
40+
}
41+
42+
} // namespace visitor
43+
} // namespace nmodl
Lines changed: 45 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,45 @@
1+
/*
2+
* Copyright 2025 EPFL.
3+
* See the top-level LICENSE file for details.
4+
*
5+
* SPDX-License-Identifier: Apache-2.0
6+
*/
7+
8+
#pragma once
9+
10+
/**
11+
* \file
12+
* \brief \copybrief nmodl::visitor::ImplicitMethodVisitor
13+
*/
14+
15+
#include "visitors/ast_visitor.hpp"
16+
17+
#include <string>
18+
#include <unordered_set>
19+
20+
21+
namespace nmodl {
22+
namespace visitor {
23+
24+
/**
25+
* \addtogroup visitor_classes
26+
* \{
27+
*/
28+
29+
/**
30+
* \class ImplicitMethodVisitor
31+
* \brief %Visitor for adding implicit method to SOLVE blocks
32+
*/
33+
class ImplicitMethodVisitor: public AstVisitor {
34+
private:
35+
std::unordered_set<std::string> derivative_block_names;
36+
37+
public:
38+
void visit_program(ast::Program& node) override;
39+
void visit_solve_block(ast::SolveBlock& node) override;
40+
};
41+
42+
/** \} */ // end of visitor_classes
43+
44+
} // namespace visitor
45+
} // namespace nmodl

test/nmodl/transpiler/unit/CMakeLists.txt

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -47,6 +47,7 @@ add_executable(
4747
visitor/defuse_analyze.cpp
4848
visitor/global_to_range.cpp
4949
visitor/implicit_argument.cpp
50+
visitor/implicit_method.cpp
5051
visitor/inline.cpp
5152
visitor/initial_block.cpp
5253
visitor/json.cpp
Lines changed: 77 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,77 @@
1+
/*
2+
* Copyright 2023 Blue Brain Project, EPFL.
3+
* See the top-level LICENSE file for details.
4+
*
5+
* SPDX-License-Identifier: Apache-2.0
6+
*/
7+
8+
#include "ast/program.hpp"
9+
#include "parser/nmodl_driver.hpp"
10+
#include "utils/test_utils.hpp"
11+
#include "visitors/implicit_method_visitor.hpp"
12+
#include "visitors/nmodl_visitor.hpp"
13+
#include "visitors/symtab_visitor.hpp"
14+
#include "visitors/visitor_utils.hpp"
15+
16+
#include <catch2/catch_test_macros.hpp>
17+
#include <catch2/matchers/catch_matchers_string.hpp>
18+
19+
using namespace nmodl;
20+
using nmodl::test_utils::reindent_text;
21+
22+
using Catch::Matchers::ContainsSubstring; // ContainsSubstring in newer Catch2
23+
24+
//=============================================================================
25+
// Implicit visitor tests
26+
//=============================================================================
27+
28+
std::string generate_mod_after_implicit_method_visitor(std::string const& text) {
29+
parser::NmodlDriver driver{};
30+
auto const ast = driver.parse_string(text);
31+
visitor::SymtabVisitor{}.visit_program(*ast);
32+
visitor::ImplicitMethodVisitor{}.visit_program(*ast);
33+
return to_nmodl(*ast);
34+
}
35+
36+
SCENARIO("Check insertion of explicit arguments to SOLVE block", "[codegen][implicit_methods]") {
37+
GIVEN("A mod file that has a SOLVE block of a derivative without an explicit METHOD") {
38+
auto const nmodl_text = R"(
39+
NEURON {
40+
SUFFIX ImplicitMethodTest
41+
}
42+
43+
BREAKPOINT {
44+
SOLVE states
45+
}
46+
47+
ASSIGNED {
48+
n
49+
}
50+
51+
DERIVATIVE states {
52+
n' = -n
53+
}
54+
)";
55+
auto const expected_text = R"(
56+
NEURON {
57+
SUFFIX ImplicitMethodTest
58+
}
59+
60+
BREAKPOINT {
61+
SOLVE states METHOD derivimplicit
62+
}
63+
64+
ASSIGNED {
65+
n
66+
}
67+
68+
DERIVATIVE states {
69+
n' = -n
70+
}
71+
)";
72+
auto const actual_text = generate_mod_after_implicit_method_visitor(nmodl_text);
73+
THEN("at_time should have nt as its first argument") {
74+
REQUIRE(reindent_text(actual_text) == reindent_text(expected_text));
75+
}
76+
}
77+
}

0 commit comments

Comments
 (0)