Skip to content

Commit 26ffb5f

Browse files
authored
Merge pull request #13 from brandmaier/codex/expose-network-classes-to-r-using-rcpp
Add Rcpp-based Network/Node bridge for R users
2 parents f4dad3c + 17f28b9 commit 26ffb5f

4 files changed

Lines changed: 302 additions & 1 deletion

File tree

DESCRIPTION

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -6,7 +6,8 @@ Maintainer: Andreas M. Brandmaier <andy@brandmaier.de>
66
Depends:
77
R (>= 3.1.0),
88
ggplot2,
9-
gridExtra
9+
gridExtra,
10+
Rcpp
1011
Suggests:
1112
knitr,
1213
rmarkdown
@@ -20,3 +21,5 @@ RoxygenNote: 7.1.0
2021
SystemRequirements: C++11
2122
VignetteBuilder: knitr
2223
NeedsCompilation: yes
24+
25+
LinkingTo: Rcpp

NAMESPACE

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -26,3 +26,22 @@ export(train.bnn)
2626
useDynLib(bnnlib)
2727
export(NetworkFactory_createFeedForwardNetwork)
2828
exportPattern("*")
29+
30+
export(bnn_create_feedforward_network)
31+
export(bnn_network_get_node)
32+
export(bnn_network_node_names)
33+
S3method(print,bnn_network_ptr)
34+
S3method(print,bnn_node_ptr)
35+
export(bnn_sequence_create)
36+
export(bnn_sequence_add)
37+
export(bnn_sequence_size)
38+
export(bnn_sequence_get_input)
39+
export(bnn_sequence_get_target)
40+
export(bnn_sequenceset_create)
41+
export(bnn_sequenceset_add_sequence)
42+
export(bnn_sequenceset_size)
43+
export(bnn_sequenceset_get_sequence)
44+
export(bnn_create_trainer)
45+
export(bnn_trainer_name)
46+
export(bnn_trainer_train)
47+
S3method(print,bnn_trainer_ptr)

R/rcpp_network.R

Lines changed: 121 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,121 @@
1+
#' Create a feedforward network object backed by a C++ pointer
2+
#' @param in_size Number of input units.
3+
#' @param hid_size Number of hidden units.
4+
#' @param out_size Number of output units.
5+
#' @return An object of class `bnn_network_ptr`.
6+
#' @export
7+
bnn_create_feedforward_network <- function(in_size, hid_size, out_size) {
8+
ptr <- .Call("bnn_create_feedforward_network", as.integer(in_size), as.integer(hid_size), as.integer(out_size))
9+
structure(list(ptr = ptr), class = "bnn_network_ptr")
10+
}
11+
12+
#' @export
13+
print.bnn_network_ptr <- function(x, ...) {
14+
n <- .Call("bnn_network_num_nodes", x$ptr)
15+
cat("<bnn_network_ptr>", n, "nodes\n")
16+
invisible(x)
17+
}
18+
19+
#' @export
20+
bnn_network_node_names <- function(network) {
21+
stopifnot(inherits(network, "bnn_network_ptr"))
22+
.Call("bnn_network_node_names", network$ptr)
23+
}
24+
25+
#' @export
26+
bnn_network_get_node <- function(network, index_one_based) {
27+
stopifnot(inherits(network, "bnn_network_ptr"))
28+
ptr <- .Call("bnn_network_get_node", network$ptr, as.integer(index_one_based))
29+
structure(list(ptr = ptr), class = "bnn_node_ptr")
30+
}
31+
32+
#' @export
33+
print.bnn_node_ptr <- function(x, ...) {
34+
name <- .Call("bnn_node_name", x$ptr)
35+
nin <- .Call("bnn_node_num_incoming", x$ptr)
36+
nout <- .Call("bnn_node_num_outgoing", x$ptr)
37+
cat("<bnn_node_ptr>", name, sprintf("(in=%d, out=%d)", nin, nout), "\n")
38+
invisible(x)
39+
}
40+
41+
#' @export
42+
bnn_sequence_create <- function() {
43+
structure(list(ptr = .Call("bnn_sequence_create")), class = "bnn_sequence_ptr")
44+
}
45+
46+
#' @export
47+
bnn_sequence_add <- function(sequence, input, target) {
48+
stopifnot(inherits(sequence, "bnn_sequence_ptr"))
49+
.Call("bnn_sequence_add", sequence$ptr, as.numeric(input), as.numeric(target))
50+
invisible(sequence)
51+
}
52+
53+
#' @export
54+
bnn_sequence_size <- function(sequence) {
55+
stopifnot(inherits(sequence, "bnn_sequence_ptr"))
56+
.Call("bnn_sequence_size", sequence$ptr)
57+
}
58+
59+
#' @export
60+
bnn_sequence_get_input <- function(sequence, index_one_based) {
61+
stopifnot(inherits(sequence, "bnn_sequence_ptr"))
62+
.Call("bnn_sequence_get_input", sequence$ptr, as.integer(index_one_based))
63+
}
64+
65+
#' @export
66+
bnn_sequence_get_target <- function(sequence, index_one_based) {
67+
stopifnot(inherits(sequence, "bnn_sequence_ptr"))
68+
.Call("bnn_sequence_get_target", sequence$ptr, as.integer(index_one_based))
69+
}
70+
71+
#' @export
72+
bnn_sequenceset_create <- function() {
73+
structure(list(ptr = .Call("bnn_sequenceset_create")), class = "bnn_sequenceset_ptr")
74+
}
75+
76+
#' @export
77+
bnn_sequenceset_add_sequence <- function(sequenceset, sequence) {
78+
stopifnot(inherits(sequenceset, "bnn_sequenceset_ptr"), inherits(sequence, "bnn_sequence_ptr"))
79+
.Call("bnn_sequenceset_add_sequence", sequenceset$ptr, sequence$ptr)
80+
invisible(sequenceset)
81+
}
82+
83+
#' @export
84+
bnn_sequenceset_size <- function(sequenceset) {
85+
stopifnot(inherits(sequenceset, "bnn_sequenceset_ptr"))
86+
.Call("bnn_sequenceset_size", sequenceset$ptr)
87+
}
88+
89+
#' @export
90+
bnn_sequenceset_get_sequence <- function(sequenceset, index_one_based) {
91+
stopifnot(inherits(sequenceset, "bnn_sequenceset_ptr"))
92+
ptr <- .Call("bnn_sequenceset_get_sequence", sequenceset$ptr, as.integer(index_one_based))
93+
structure(list(ptr = ptr), class = "bnn_sequence_ptr")
94+
}
95+
96+
#' @export
97+
bnn_create_trainer <- function(network, trainer_type = c("backprop", "adam", "rmsprop", "rprop", "myrprop")) {
98+
stopifnot(inherits(network, "bnn_network_ptr"))
99+
trainer_type <- match.arg(trainer_type)
100+
ptr <- .Call("bnn_create_trainer", trainer_type, network$ptr)
101+
structure(list(ptr = ptr, type = trainer_type), class = "bnn_trainer_ptr")
102+
}
103+
104+
#' @export
105+
bnn_trainer_name <- function(trainer) {
106+
stopifnot(inherits(trainer, "bnn_trainer_ptr"))
107+
.Call("bnn_trainer_name", trainer$ptr)
108+
}
109+
110+
#' @export
111+
bnn_trainer_train <- function(trainer, sequenceset, iterations) {
112+
stopifnot(inherits(trainer, "bnn_trainer_ptr"), inherits(sequenceset, "bnn_sequenceset_ptr"))
113+
.Call("bnn_trainer_train", trainer$ptr, sequenceset$ptr, as.integer(iterations))
114+
invisible(trainer)
115+
}
116+
117+
#' @export
118+
print.bnn_trainer_ptr <- function(x, ...) {
119+
cat("<bnn_trainer_ptr>", .Call("bnn_trainer_name", x$ptr), "\n")
120+
invisible(x)
121+
}

src/rcpp_network.cpp

Lines changed: 158 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,158 @@
1+
#include <Rcpp.h>
2+
#include "NetworkFactory.h"
3+
#include "Network.h"
4+
#include "nodes/Node.h"
5+
#include "Sequence.h"
6+
#include "SequenceSet.h"
7+
#include "trainer/Trainer.h"
8+
#include "trainer/BackpropTrainer.h"
9+
#include "trainer/ADAMTrainer.h"
10+
#include "trainer/RMSPropTrainer.h"
11+
#include "trainer/RPropTrainer.h"
12+
#include "trainer/MyRPropTrainer.h"
13+
14+
using namespace Rcpp;
15+
16+
namespace {
17+
SEXP make_node_xptr(Node* node) {
18+
XPtr<Node> ptr(node, false);
19+
return ptr;
20+
}
21+
22+
SEXP make_sequence_xptr(Sequence* sequence) {
23+
XPtr<Sequence> ptr(sequence, true);
24+
return ptr;
25+
}
26+
27+
SEXP make_trainer_xptr(Trainer* trainer) {
28+
XPtr<Trainer> ptr(trainer, true);
29+
return ptr;
30+
}
31+
}
32+
33+
extern "C" SEXP bnn_create_feedforward_network(SEXP in_sizeSEXP, SEXP hid_sizeSEXP, SEXP out_sizeSEXP) {
34+
int in_size = as<int>(in_sizeSEXP);
35+
int hid_size = as<int>(hid_sizeSEXP);
36+
int out_size = as<int>(out_sizeSEXP);
37+
if (in_size <= 0 || hid_size <= 0 || out_size <= 0) stop("All layer sizes must be positive integers.");
38+
Network* network = NetworkFactory::createFeedForwardNetwork((unsigned int)in_size, (unsigned int)hid_size, (unsigned int)out_size);
39+
XPtr<Network> ptr(network, true);
40+
return ptr;
41+
}
42+
43+
extern "C" SEXP bnn_network_num_nodes(SEXP network_xptr) {
44+
XPtr<Network> net(network_xptr);
45+
return wrap((int)net->get_num_nodes());
46+
}
47+
48+
extern "C" SEXP bnn_network_node_names(SEXP network_xptr) {
49+
XPtr<Network> net(network_xptr);
50+
return wrap(net->get_node_names());
51+
}
52+
53+
extern "C" SEXP bnn_network_get_node(SEXP network_xptr, SEXP index_one_basedSEXP) {
54+
XPtr<Network> net(network_xptr);
55+
int idx = as<int>(index_one_basedSEXP);
56+
if (idx < 1 || idx > (int)net->nodes.size()) stop("index_one_based is out of range.");
57+
return make_node_xptr(net->nodes[(size_t)(idx - 1)]);
58+
}
59+
60+
extern "C" SEXP bnn_node_name(SEXP node_xptr) {
61+
XPtr<Node> node(node_xptr);
62+
return wrap(node->name);
63+
}
64+
65+
extern "C" SEXP bnn_node_num_incoming(SEXP node_xptr) {
66+
XPtr<Node> node(node_xptr);
67+
return wrap((int)node->get_num_incoming_connections());
68+
}
69+
70+
extern "C" SEXP bnn_node_num_outgoing(SEXP node_xptr) {
71+
XPtr<Node> node(node_xptr);
72+
return wrap((int)node->get_num_outgoing_connections());
73+
}
74+
75+
extern "C" SEXP bnn_sequence_create() {
76+
return make_sequence_xptr(new Sequence());
77+
}
78+
79+
extern "C" SEXP bnn_sequence_add(SEXP sequence_xptr, SEXP inputSEXP, SEXP targetSEXP) {
80+
XPtr<Sequence> sequence(sequence_xptr);
81+
NumericVector input = as<NumericVector>(inputSEXP);
82+
NumericVector target = as<NumericVector>(targetSEXP);
83+
std::vector<weight_t>* in = new std::vector<weight_t>(input.begin(), input.end());
84+
std::vector<weight_t>* tar = new std::vector<weight_t>(target.begin(), target.end());
85+
sequence->add(in, tar);
86+
return R_NilValue;
87+
}
88+
89+
extern "C" SEXP bnn_sequence_size(SEXP sequence_xptr) {
90+
XPtr<Sequence> sequence(sequence_xptr);
91+
return wrap((int)sequence->size());
92+
}
93+
94+
extern "C" SEXP bnn_sequence_get_input(SEXP sequence_xptr, SEXP index_one_basedSEXP) {
95+
XPtr<Sequence> sequence(sequence_xptr);
96+
int idx = as<int>(index_one_basedSEXP);
97+
if (idx < 1 || idx > (int)sequence->size()) stop("index_one_based is out of range.");
98+
return wrap(*sequence->get_input((unsigned int)(idx - 1)));
99+
}
100+
101+
extern "C" SEXP bnn_sequence_get_target(SEXP sequence_xptr, SEXP index_one_basedSEXP) {
102+
XPtr<Sequence> sequence(sequence_xptr);
103+
int idx = as<int>(index_one_basedSEXP);
104+
if (idx < 1 || idx > (int)sequence->size()) stop("index_one_based is out of range.");
105+
return wrap(*sequence->get_target((unsigned int)(idx - 1)));
106+
}
107+
108+
extern "C" SEXP bnn_sequenceset_create() {
109+
XPtr<SequenceSet> ptr(new SequenceSet(), true);
110+
return ptr;
111+
}
112+
113+
extern "C" SEXP bnn_sequenceset_add_sequence(SEXP sequenceset_xptr, SEXP sequence_xptr) {
114+
XPtr<SequenceSet> set(sequenceset_xptr);
115+
XPtr<Sequence> seq(sequence_xptr);
116+
set->add_copy_of_sequence(seq.get());
117+
return R_NilValue;
118+
}
119+
120+
extern "C" SEXP bnn_sequenceset_size(SEXP sequenceset_xptr) {
121+
XPtr<SequenceSet> set(sequenceset_xptr);
122+
return wrap((int)set->size());
123+
}
124+
125+
extern "C" SEXP bnn_sequenceset_get_sequence(SEXP sequenceset_xptr, SEXP index_one_basedSEXP) {
126+
XPtr<SequenceSet> set(sequenceset_xptr);
127+
int idx = as<int>(index_one_basedSEXP);
128+
if (idx < 1 || idx > (int)set->size()) stop("index_one_based is out of range.");
129+
XPtr<Sequence> ptr(new Sequence(*set->get((unsigned int)(idx - 1))), true);
130+
return ptr;
131+
}
132+
133+
extern "C" SEXP bnn_create_trainer(SEXP trainer_typeSEXP, SEXP network_xptr) {
134+
std::string trainer_type = as<std::string>(trainer_typeSEXP);
135+
XPtr<Network> network(network_xptr);
136+
137+
if (trainer_type == "backprop") return make_trainer_xptr(new BackpropTrainer(network.get()));
138+
if (trainer_type == "adam") return make_trainer_xptr(new ADAMTrainer(network.get()));
139+
if (trainer_type == "rmsprop") return make_trainer_xptr(new RMSPropTrainer(network.get()));
140+
if (trainer_type == "rprop") return make_trainer_xptr(new RPropTrainer(network.get()));
141+
if (trainer_type == "myrprop") return make_trainer_xptr(new MyRPropTrainer(network.get()));
142+
143+
stop("Unsupported trainer_type. Use one of: backprop, adam, rmsprop, rprop, myrprop.");
144+
}
145+
146+
extern "C" SEXP bnn_trainer_name(SEXP trainer_xptr) {
147+
XPtr<Trainer> trainer(trainer_xptr);
148+
return wrap(trainer->get_name());
149+
}
150+
151+
extern "C" SEXP bnn_trainer_train(SEXP trainer_xptr, SEXP sequenceset_xptr, SEXP iterationsSEXP) {
152+
XPtr<Trainer> trainer(trainer_xptr);
153+
XPtr<SequenceSet> set(sequenceset_xptr);
154+
int iterations = as<int>(iterationsSEXP);
155+
if (iterations <= 0) stop("iterations must be a positive integer.");
156+
trainer->train(set.get(), (unsigned int)iterations);
157+
return R_NilValue;
158+
}

0 commit comments

Comments
 (0)