IgANet
IGAnets - Isogeometric Analysis Networks
Loading...
Searching...
No Matches
generator.hpp
Go to the documentation of this file.
1
15#pragma once
16
17#include <any>
18#include <iostream>
19#include <vector>
20
21#include <iganet/core/core.hpp>
24#include <iganet/utils/zip.hpp>
25
26namespace iganet {
27
28// clang-format off
30enum class nn_init : short_t {
31 constant = 0,
32 normal = 1,
33 uniform = 2,
34 kaiming_normal = 3,
35 kaiming_uniform = 4,
36 xavier_normal = 5,
37 xavier_uniform = 6,
38};
39// clang-format on
40
50template <typename real_t>
51class IgANetGeneratorImpl : public torch::nn::Module {
52public:
55
61 const std::vector<int64_t> &layers,
62 const std::vector<std::vector<std::any>> &activations,
63 Options<real_t> options = Options<real_t>{}) {
64 assert(layers.size() == activations.size() + 1);
65
66 // Generate vector of linear layers and register them as layer[i]
67 for (auto i = 0; i < layers.size() - 1; ++i) {
68 layers_.emplace_back(
69 register_module("layer[" + std::to_string(i) + "]",
70 torch::nn::Linear(layers[i], layers[i + 1])));
71 layers_.back()->to(options.device(), options.dtype(), true);
72
73 torch::nn::init::xavier_uniform_(layers_.back()->weight);
74 torch::nn::init::constant_(layers_.back()->bias, 0.0);
75 }
76
77 // Generate vector of activation functions
78 for (const auto &a : activations)
79 switch (std::any_cast<activation>(a[0])) {
80 // No activation function
82 switch (a.size()) {
83 case 1:
84 activations_.emplace_back(new None{});
85 break;
86 default:
87 throw std::runtime_error("Invalid number of parameters");
88 }
89 break;
90
91 // Batch Normalization
93 switch (a.size()) {
94 case 8:
95 activations_.emplace_back(new BatchNorm{
96 std::any_cast<torch::Tensor>(a[1]),
97 std::any_cast<torch::Tensor>(a[2]),
98 std::any_cast<torch::Tensor>(a[3]),
99 std::any_cast<torch::Tensor>(a[4]), std::any_cast<double>(a[5]),
100 std::any_cast<double>(a[6]), std::any_cast<bool>(a[7])});
101 break;
102 case 7:
103 activations_.emplace_back(new BatchNorm{
104 std::any_cast<torch::Tensor>(a[1]),
105 std::any_cast<torch::Tensor>(a[2]),
106 std::any_cast<torch::Tensor>(a[3]),
107 std::any_cast<torch::Tensor>(a[4]), std::any_cast<double>(a[5]),
108 std::any_cast<double>(a[6])});
109 break;
110 case 4:
111 activations_.emplace_back(new BatchNorm{
112 std::any_cast<torch::Tensor>(a[1]),
113 std::any_cast<torch::Tensor>(a[2]),
114 std::any_cast<torch::nn::functional::BatchNormFuncOptions>(
115 a[3])});
116 break;
117 case 3:
118 activations_.emplace_back(
119 new BatchNorm{std::any_cast<torch::Tensor>(a[1]),
120 std::any_cast<torch::Tensor>(a[2])});
121 break;
122 default:
123 throw std::runtime_error("Invalid number of parameters");
124 }
125 break;
126
127 // CELU
128 case activation::celu:
129 switch (a.size()) {
130 case 3:
131 activations_.emplace_back(
132 new CELU{std::any_cast<double>(a[1]), std::any_cast<bool>(a[2])});
133 break;
134 case 2:
135 try {
136 activations_.emplace_back(new CELU{
137 std::any_cast<torch::nn::functional::CELUFuncOptions>(a[1])});
138 } catch (...) {
139 activations_.emplace_back(new CELU{std::any_cast<double>(a[1])});
140 }
141 break;
142 case 1:
143 activations_.emplace_back(new CELU{});
144 break;
145 default:
146 throw std::runtime_error("Invalid number of parameters");
147 }
148 break;
149
150 // ELU
151 case activation::elu:
152 switch (a.size()) {
153 case 3:
154 activations_.emplace_back(
155 new ELU{std::any_cast<double>(a[1]), std::any_cast<bool>(a[2])});
156 break;
157 case 2:
158 try {
159 activations_.emplace_back(new ELU{
160 std::any_cast<torch::nn::functional::ELUFuncOptions>(a[1])});
161 } catch (...) {
162 activations_.emplace_back(new ELU{std::any_cast<double>(a[1])});
163 }
164 break;
165 case 1:
166 activations_.emplace_back(new ELU{});
167 break;
168 default:
169 throw std::runtime_error("Invalid number of parameters");
170 }
171 break;
172
173 // GELU
174 case activation::gelu:
175 switch (a.size()) {
176 case 1:
177 activations_.emplace_back(new GELU{});
178 break;
179 default:
180 throw std::runtime_error("Invalid number of parameters");
181 }
182 break;
183
184 // GLU
185 case activation::glu:
186 switch (a.size()) {
187 case 2:
188 try {
189 activations_.emplace_back(new GLU{
190 std::any_cast<torch::nn::functional::GLUFuncOptions>(a[1])});
191 } catch (...) {
192 activations_.emplace_back(new GLU{std::any_cast<int64_t>(a[1])});
193 }
194 break;
195 case 1:
196 activations_.emplace_back(new GLU{});
197 break;
198 default:
199 throw std::runtime_error("Invalid number of parameters");
200 }
201 break;
202
203 // Group Normalization
205 switch (a.size()) {
206 case 5:
207 activations_.emplace_back(new GroupNorm{
208 std::any_cast<int64_t>(a[1]), std::any_cast<torch::Tensor>(a[2]),
209 std::any_cast<torch::Tensor>(a[3]), std::any_cast<double>(a[4])});
210 break;
211 case 2:
212 try {
213 activations_.emplace_back(new GroupNorm{
214 std::any_cast<torch::nn::functional::GroupNormFuncOptions>(
215 a[1])});
216 } catch (...) {
217 activations_.emplace_back(
218 new GroupNorm{std::any_cast<int64_t>(a[1])});
219 }
220 break;
221 default:
222 throw std::runtime_error("Invalid number of parameters");
223 }
224 break;
225
226 // Gumbel-Softmax
228 switch (a.size()) {
229 case 4:
230 activations_.emplace_back(new GumbelSoftmax{
231 std::any_cast<double>(a[1]), std::any_cast<int>(a[2]),
232 std::any_cast<bool>(a[3])});
233 break;
234 case 2:
235 activations_.emplace_back(new GumbelSoftmax{
236 std::any_cast<torch::nn::functional::GumbelSoftmaxFuncOptions>(
237 a[1])});
238 break;
239 case 1:
240 activations_.emplace_back(new GumbelSoftmax{});
241 break;
242 default:
243 throw std::runtime_error("Invalid number of parameters");
244 }
245 break;
246
247 // Hard shrinkish
249 switch (a.size()) {
250 case 2:
251 try {
252 activations_.emplace_back(new Hardshrink{
253 std::any_cast<torch::nn::functional::HardshrinkFuncOptions>(
254 a[1])});
255 } catch (...) {
256 activations_.emplace_back(
257 new Hardshrink{std::any_cast<double>(a[1])});
258 }
259 break;
260 case 1:
261 activations_.emplace_back(new Hardshrink{});
262 break;
263 default:
264 throw std::runtime_error("Invalid number of parameters");
265 }
266 break;
267
268 // Hardsigmoid
270 switch (a.size()) {
271 case 1:
272 activations_.emplace_back(new Hardsigmoid{});
273 break;
274 default:
275 throw std::runtime_error("Invalid number of parameters");
276 }
277 break;
278
279 // Hardswish
281 switch (a.size()) {
282 case 1:
283 activations_.emplace_back(new Hardswish{});
284 break;
285 default:
286 throw std::runtime_error("Invalid number of parameters");
287 }
288 break;
289
290 // Hardtanh
292 switch (a.size()) {
293 case 4:
294 activations_.emplace_back(new Hardtanh{std::any_cast<double>(a[1]),
295 std::any_cast<double>(a[2]),
296 std::any_cast<bool>(a[3])});
297 break;
298 case 3:
299 activations_.emplace_back(new Hardtanh{std::any_cast<double>(a[1]),
300 std::any_cast<double>(a[2])});
301 break;
302 case 2:
303 activations_.emplace_back(new Hardtanh{
304 std::any_cast<torch::nn::functional::HardtanhFuncOptions>(a[1])});
305 break;
306 case 1:
307 activations_.emplace_back(new Hardtanh{});
308 break;
309 default:
310 throw std::runtime_error("Invalid number of parameters");
311 }
312 break;
313
314 // Instance Normalization
316 switch (a.size()) {
317 case 8:
318 activations_.emplace_back(new InstanceNorm{
319 std::any_cast<torch::Tensor>(a[1]),
320 std::any_cast<torch::Tensor>(a[2]),
321 std::any_cast<torch::Tensor>(a[3]),
322 std::any_cast<torch::Tensor>(a[4]), std::any_cast<double>(a[5]),
323 std::any_cast<double>(a[6]), std::any_cast<bool>(a[7])});
324 break;
325 case 7:
326 activations_.emplace_back(new InstanceNorm{
327 std::any_cast<torch::Tensor>(a[1]),
328 std::any_cast<torch::Tensor>(a[2]),
329 std::any_cast<torch::Tensor>(a[3]),
330 std::any_cast<torch::Tensor>(a[4]), std::any_cast<double>(a[5]),
331 std::any_cast<double>(a[6])});
332 break;
333 case 2:
334 activations_.emplace_back(new InstanceNorm{
335 std::any_cast<torch::nn::functional::InstanceNormFuncOptions>(
336 a[1])});
337 break;
338 case 1:
339 activations_.emplace_back(new InstanceNorm{});
340 break;
341 default:
342 throw std::runtime_error("Invalid number of parameters");
343 }
344 break;
345
346 // Layer Normalization
348 switch (a.size()) {
349 case 5:
350 activations_.emplace_back(new LayerNorm{
351 std::any_cast<std::vector<int64_t>>(a[1]),
352 std::any_cast<torch::Tensor>(a[2]),
353 std::any_cast<torch::Tensor>(a[3]), std::any_cast<double>(a[4])});
354 break;
355 case 2:
356 try {
357 activations_.emplace_back(new LayerNorm{
358 std::any_cast<torch::nn::functional::LayerNormFuncOptions>(
359 a[1])});
360 } catch (...) {
361 activations_.emplace_back(
362 new LayerNorm{std::any_cast<std::vector<int64_t>>(a[1])});
363 }
364 break;
365 default:
366 throw std::runtime_error("Invalid number of parameters");
367 }
368 break;
369
370 // Leaky ReLU
372 switch (a.size()) {
373 case 3:
374 activations_.emplace_back(new LeakyReLU{std::any_cast<double>(a[1]),
375 std::any_cast<bool>(a[2])});
376 break;
377 case 2:
378 try {
379 activations_.emplace_back(new LeakyReLU{
380 std::any_cast<torch::nn::functional::LeakyReLUFuncOptions>(
381 a[1])});
382 } catch (...) {
383 activations_.emplace_back(
384 new LeakyReLU{std::any_cast<double>(a[1])});
385 }
386 break;
387 case 1:
388 activations_.emplace_back(new LeakyReLU{});
389 break;
390 default:
391 throw std::runtime_error("Invalid number of parameters");
392 }
393 break;
394
395 // Local response Normalization
397 switch (a.size()) {
398 case 5:
399 activations_.emplace_back(new LocalResponseNorm{
400 std::any_cast<int64_t>(a[1]), std::any_cast<double>(a[2]),
401 std::any_cast<double>(a[3]), std::any_cast<double>(a[4])});
402 break;
403 case 2:
404 try {
405 activations_.emplace_back(new LocalResponseNorm{std::any_cast<
406 torch::nn::functional::LocalResponseNormFuncOptions>(a[1])});
407 } catch (...) {
408 activations_.emplace_back(
409 new LocalResponseNorm{std::any_cast<int64_t>(a[1])});
410 }
411 break;
412 default:
413 throw std::runtime_error("Invalid number of parameters");
414 }
415 break;
416
417 // LogSigmoid
419 switch (a.size()) {
420 case 1:
421 activations_.emplace_back(new LogSigmoid{});
422 break;
423 default:
424 throw std::runtime_error("Invalid number of parameters");
425 }
426 break;
427
428 // LogSoftmax
430 switch (a.size()) {
431 case 2:
432 try {
433 activations_.emplace_back(new LogSoftmax{
434 std::any_cast<torch::nn::functional::LogSoftmaxFuncOptions>(
435 a[1])});
436 } catch (...) {
437 activations_.emplace_back(
438 new LogSoftmax{std::any_cast<int64_t>(a[1])});
439 }
440 break;
441 default:
442 throw std::runtime_error("Invalid number of parameters");
443 }
444 break;
445
446 // Mish
447 case activation::mish:
448 switch (a.size()) {
449 case 1:
450 activations_.emplace_back(new Mish{});
451 break;
452 default:
453 throw std::runtime_error("Invalid number of parameters");
454 }
455 break;
456
457 // Lp Normalization
459 switch (a.size()) {
460 case 4:
461 activations_.emplace_back(new Normalize{
462 std::any_cast<double>(a[1]), std::any_cast<double>(a[2]),
463 std::any_cast<int64_t>(a[3])});
464 break;
465 case 2:
466 activations_.emplace_back(new Normalize{
467 std::any_cast<torch::nn::functional::NormalizeFuncOptions>(
468 a[1])});
469 break;
470 case 1:
471 activations_.emplace_back(new Normalize{});
472 break;
473 default:
474 throw std::runtime_error("Invalid number of parameters");
475 }
476 break;
477
478 // PReLU
480 switch (a.size()) {
481 case 2:
482 activations_.emplace_back(
483 new PReLU{std::any_cast<torch::Tensor>(a[1])});
484 break;
485 default:
486 throw std::runtime_error("Invalid number of parameters");
487 }
488 break;
489
490 // ReLU
491 case activation::relu:
492 switch (a.size()) {
493 case 2:
494 try {
495 activations_.emplace_back(new ReLU{
496 std::any_cast<torch::nn::functional::ReLUFuncOptions>(a[1])});
497 } catch (...) {
498 activations_.emplace_back(new ReLU{std::any_cast<bool>(a[1])});
499 }
500 break;
501 case 1:
502 activations_.emplace_back(new ReLU{});
503 break;
504 default:
505 throw std::runtime_error("Invalid number of parameters");
506 }
507 break;
508
509 // Relu6
511 switch (a.size()) {
512 case 2:
513 try {
514 activations_.emplace_back(new ReLU6{
515 std::any_cast<torch::nn::functional::ReLU6FuncOptions>(a[1])});
516 } catch (...) {
517 activations_.emplace_back(new ReLU6{std::any_cast<bool>(a[1])});
518 }
519 break;
520 case 1:
521 activations_.emplace_back(new ReLU6{});
522 break;
523 default:
524 throw std::runtime_error("Invalid number of parameters");
525 }
526 break;
527
528 // Randomized ReLU
530 switch (a.size()) {
531 case 4:
532 activations_.emplace_back(new RReLU{std::any_cast<double>(a[1]),
533 std::any_cast<double>(a[2]),
534 std::any_cast<bool>(a[3])});
535 break;
536 case 3:
537 activations_.emplace_back(new RReLU{std::any_cast<double>(a[1]),
538 std::any_cast<double>(a[2])});
539 break;
540 case 2:
541 activations_.emplace_back(new RReLU{
542 std::any_cast<torch::nn::functional::RReLUFuncOptions>(a[1])});
543 break;
544 case 1:
545 activations_.emplace_back(new RReLU{});
546 break;
547 default:
548 throw std::runtime_error("Invalid number of parameters");
549 }
550 break;
551
552 // SELU
553 case activation::selu:
554 switch (a.size()) {
555 case 2:
556 try {
557 activations_.emplace_back(new SELU{
558 std::any_cast<torch::nn::functional::SELUFuncOptions>(a[1])});
559 } catch (...) {
560 activations_.emplace_back(new SELU{std::any_cast<bool>(a[1])});
561 }
562 break;
563 case 1:
564 activations_.emplace_back(new SELU{});
565 break;
566 default:
567 throw std::runtime_error("Invalid number of parameters");
568 }
569 break;
570
571 // Sigmoid
573 switch (a.size()) {
574 case 1:
575 activations_.emplace_back(new Sigmoid{});
576 break;
577 default:
578 throw std::runtime_error("Invalid number of parameters");
579 }
580 break;
581
582 // SiLU
583 case activation::silu:
584 switch (a.size()) {
585 case 1:
586 activations_.emplace_back(new SiLU{});
587 break;
588 default:
589 throw std::runtime_error("Invalid number of parameters");
590 }
591 break;
592
593 // Softmax
595 switch (a.size()) {
596 case 2:
597 try {
598 activations_.emplace_back(new Softmax{
599 std::any_cast<torch::nn::functional::SoftmaxFuncOptions>(
600 a[1])});
601 } catch (...) {
602 activations_.emplace_back(
603 new Softmax{std::any_cast<int64_t>(a[1])});
604 }
605 break;
606 default:
607 throw std::runtime_error("Invalid number of parameters");
608 }
609 break;
610
611 // Softmin
613 switch (a.size()) {
614 case 2:
615 try {
616 activations_.emplace_back(new Softmin{
617 std::any_cast<torch::nn::functional::SoftminFuncOptions>(
618 a[1])});
619 } catch (...) {
620 activations_.emplace_back(
621 new Softmin{std::any_cast<int64_t>(a[1])});
622 }
623 break;
624 default:
625 throw std::runtime_error("Invalid number of parameters");
626 }
627 break;
628
629 // Softplus
631 switch (a.size()) {
632 case 3:
633 activations_.emplace_back(new Softplus{std::any_cast<double>(a[1]),
634 std::any_cast<double>(a[2])});
635 break;
636 case 2:
637 activations_.emplace_back(new Softplus{
638 std::any_cast<torch::nn::functional::SoftplusFuncOptions>(a[1])});
639 break;
640 case 1:
641 activations_.emplace_back(new Softplus{});
642 break;
643 default:
644 throw std::runtime_error("Invalid number of parameters");
645 }
646 break;
647
648 // Softshrink
650 switch (a.size()) {
651 case 2:
652 try {
653 activations_.emplace_back(new Softshrink{
654 std::any_cast<torch::nn::functional::SoftshrinkFuncOptions>(
655 a[1])});
656 } catch (...) {
657 activations_.emplace_back(
658 new Softshrink{std::any_cast<double>(a[1])});
659 }
660 break;
661 case 1:
662 activations_.emplace_back(new Softshrink{});
663 break;
664 default:
665 throw std::runtime_error("Invalid number of parameters");
666 }
667 break;
668
669 // Softsign
671 switch (a.size()) {
672 case 1:
673 activations_.emplace_back(new Softsign{});
674 break;
675 default:
676 throw std::runtime_error("Invalid number of parameters");
677 }
678 break;
679
680 // Tanh
681 case activation::tanh:
682 switch (a.size()) {
683 case 1:
684 activations_.emplace_back(new Tanh{});
685 break;
686 default:
687 throw std::runtime_error("Invalid number of parameters");
688 }
689 break;
690
691 // Tanhshrink
693 switch (a.size()) {
694 case 1:
695 activations_.emplace_back(new Tanhshrink{});
696 break;
697 default:
698 throw std::runtime_error("Invalid number of parameters");
699 }
700 break;
701
702 // Threshold
704 switch (a.size()) {
705 case 4:
706 activations_.emplace_back(new Threshold{std::any_cast<double>(a[1]),
707 std::any_cast<double>(a[2]),
708 std::any_cast<bool>(a[3])});
709 break;
710 case 3:
711 activations_.emplace_back(new Threshold{std::any_cast<double>(a[1]),
712 std::any_cast<double>(a[2])});
713 break;
714 case 2:
715 activations_.emplace_back(new Threshold{
716 std::any_cast<torch::nn::functional::ThresholdFuncOptions>(
717 a[1])});
718 break;
719 default:
720 throw std::runtime_error("Invalid number of parameters");
721 }
722 break;
723
724 default:
725 throw std::runtime_error("Invalid activation function");
726 }
727 }
728
732 torch::Tensor forward(torch::Tensor x) {
733 torch::Tensor x_in = x.clone();
734
735 // Standard feed-forward neural network
736 for (auto [layer, activation] : utils::zip(layers_, activations_))
737 x = activation->apply(layer->forward(x));
738
739 return x;
740 }
741
746 inline torch::serialize::OutputArchive &
747 write(torch::serialize::OutputArchive &archive,
748 const std::string &key = "iganet") const {
749 assert(layers_.size() == activations_.size());
750
751 archive.write(key + ".layers",
752 torch::full({1}, static_cast<int64_t>(layers_.size())));
753 for (std::size_t i = 0; i < layers_.size(); ++i) {
754 archive.write(
755 key + ".layer[" + std::to_string(i) + "].in_features",
756 torch::full({1}, (int64_t)layers_[i]->options.in_features()));
757 archive.write(
758 key + ".layer[" + std::to_string(i) + "].outputs_features",
759 torch::full({1}, (int64_t)layers_[i]->options.out_features()));
760 archive.write(key + ".layer[" + std::to_string(i) + "].bias",
761 torch::full({1}, (int64_t)layers_[i]->options.bias()));
762
763 activations_[i]->write(archive, key + ".layer[" + std::to_string(i) +
764 "].activation");
765 }
766
767 return archive;
768 }
769
774 inline torch::serialize::InputArchive &
775 read(torch::serialize::InputArchive &archive,
776 const std::string &key = "iganet") {
777 torch::Tensor layers, in_features, outputs_features, bias, activation;
778
779 auto options = iganet::Options<real_t>{};
780
781 archive.read(key + ".layers", layers);
782 for (int64_t i = 0; i < layers.item<int64_t>(); ++i) {
783 archive.read(key + ".layer[" + std::to_string(i) + "].in_features",
784 in_features);
785 archive.read(key + ".layer[" + std::to_string(i) + "].outputs_features",
786 outputs_features);
787 archive.read(key + ".layer[" + std::to_string(i) + "].bias", bias);
788 layers_.emplace_back(register_module(
789 "layer[" + std::to_string(i) + "]",
790 torch::nn::Linear(
791 torch::nn::LinearOptions(in_features.item<int64_t>(),
792 outputs_features.item<int64_t>())
793 .bias(bias.item<bool>()))));
794 layers_.back()->to(options.device(), options.dtype(), true);
795
796 archive.read(key + ".layer[" + std::to_string(i) + "].activation.type",
797 activation);
798 switch (static_cast<enum activation>(activation.item<int64_t>())) {
799 case activation::none:
800 activations_.emplace_back(new None{});
801 break;
803 activations_.emplace_back(
804 new BatchNorm{torch::Tensor{}, torch::Tensor{}});
805 break;
806 case activation::celu:
807 activations_.emplace_back(new CELU{});
808 break;
809 case activation::elu:
810 activations_.emplace_back(new ELU{});
811 break;
812 case activation::gelu:
813 activations_.emplace_back(new GELU{});
814 break;
815 case activation::glu:
816 activations_.emplace_back(new GLU{});
817 break;
819 activations_.emplace_back(new GroupNorm{0});
820 break;
822 activations_.emplace_back(new GumbelSoftmax{});
823 break;
825 activations_.emplace_back(new Hardshrink{});
826 break;
828 activations_.emplace_back(new Hardsigmoid{});
829 break;
831 activations_.emplace_back(new Hardswish{});
832 break;
834 activations_.emplace_back(new Hardtanh{});
835 break;
837 activations_.emplace_back(new InstanceNorm{});
838 break;
840 activations_.emplace_back(new LayerNorm{{}});
841 break;
843 activations_.emplace_back(new LeakyReLU{});
844 break;
846 activations_.emplace_back(new LocalResponseNorm{0});
847 break;
849 activations_.emplace_back(new LogSigmoid{});
850 break;
852 activations_.emplace_back(new LogSoftmax{0});
853 break;
854 case activation::mish:
855 activations_.emplace_back(new Mish{});
856 break;
858 activations_.emplace_back(new Normalize{0, 0, 0});
859 break;
861 activations_.emplace_back(new PReLU{torch::Tensor{}});
862 break;
863 case activation::relu:
864 activations_.emplace_back(new ReLU{});
865 break;
867 activations_.emplace_back(new ReLU6{});
868 break;
870 activations_.emplace_back(new RReLU{});
871 break;
872 case activation::selu:
873 activations_.emplace_back(new SELU{});
874 break;
876 activations_.emplace_back(new Sigmoid{});
877 break;
878 case activation::silu:
879 activations_.emplace_back(new SiLU{});
880 break;
882 activations_.emplace_back(new Softmax{0});
883 break;
885 activations_.emplace_back(new Softmin{0});
886 break;
888 activations_.emplace_back(new Softplus{});
889 break;
891 activations_.emplace_back(new Softshrink{});
892 break;
894 activations_.emplace_back(new Softsign{});
895 break;
896 case activation::tanh:
897 activations_.emplace_back(new Tanh{});
898 break;
900 activations_.emplace_back(new Tanhshrink{});
901 break;
903 activations_.emplace_back(new Threshold{0, 0});
904 break;
905 default:
906 throw std::runtime_error("Invalid activation function");
907 }
908 activations_.back()->read(archive, key + ".layer[" + std::to_string(i) +
909 "].activation");
910 }
911 return archive;
912 }
913
916 inline void pretty_print(std::ostream &os) const noexcept override {
917 os << "(\n";
918
919 int i = 0;
920 for (const auto &activation : activations_)
921 os << "activation[" << i++ << "] = " << *activation << "\n";
922 os << ")\n";
923 }
924
925private:
927 std::vector<torch::nn::Linear> layers_;
928
930 std::vector<std::unique_ptr<iganet::ActivationFunction>> activations_;
931};
932
938template <typename real_t>
940 : public torch::nn::ModuleHolder<IgANetGeneratorImpl<real_t>> {
941
942public:
943 using torch::nn::ModuleHolder<IgANetGeneratorImpl<real_t>>::ModuleHolder;
945};
946
947} // namespace iganet
Activation functions.
Batch Normalization as described in the paper.
Definition activation.hpp:159
Continuously Differentiable Exponential Linear Units activation function.
Definition activation.hpp:316
Exponential Linear Units activation function.
Definition activation.hpp:409
Gaussian Error Linear Units activation function.
Definition activation.hpp:502
Grated Linear Units activation function.
Definition activation.hpp:562
Group Normalization over a mini-batch of inputs as described in the paper Group Normalization,...
Definition activation.hpp:642
Gumbel-Softmax distribution activation function.
Definition activation.hpp:745
Hard shrinkish activation function.
Definition activation.hpp:841
Hardsigmoid activation function.
Definition activation.hpp:933
Hardswish activation function.
Definition activation.hpp:996
Hardtanh activation function.
Definition activation.hpp:1058
IgANetGenerator.
Definition generator.hpp:940
IgANetGeneratorImpl.
Definition generator.hpp:51
torch::serialize::InputArchive & read(torch::serialize::InputArchive &archive, const std::string &key="iganet")
Reads the IgANet from a torch::serialize::InputArchive object.
Definition generator.hpp:775
IgANetGeneratorImpl()=default
Default constructor.
void pretty_print(std::ostream &os) const noexcept override
Provides the pretty_print operation.
Definition generator.hpp:916
IgANetGeneratorImpl(const std::vector< int64_t > &layers, const std::vector< std::vector< std::any > > &activations, Options< real_t > options=Options< real_t >{})
Constructor.
Definition generator.hpp:60
std::vector< std::unique_ptr< iganet::ActivationFunction > > activations_
Vector of activation functions.
Definition generator.hpp:930
torch::Tensor forward(torch::Tensor x)
Forward evaluation.
Definition generator.hpp:732
std::vector< torch::nn::Linear > layers_
Vector of linear layers.
Definition generator.hpp:927
torch::serialize::OutputArchive & write(torch::serialize::OutputArchive &archive, const std::string &key="iganet") const
Writes the IgANet into a torch::serialize::OutputArchive object.
Definition generator.hpp:747
Instance Normalization as described in the paper.
Definition activation.hpp:1159
Layer Normalization as described in the paper.
Definition activation.hpp:1288
Leaky ReLU activation function.
Definition activation.hpp:1402
Local response Normalization.
Definition activation.hpp:1492
LogSigmoid activation function.
Definition activation.hpp:1605
LogSoftmax activation function.
Definition activation.hpp:1664
Mish activation function.
Definition activation.hpp:1745
No-op activation function.
Definition activation.hpp:107
Lp Normalization.
Definition activation.hpp:1797
The Options class handles the automated determination of dtype from the template argument and the sel...
Definition options.hpp:47
PReLU activation function.
Definition activation.hpp:1888
Randomized ReLU activation function.
Definition activation.hpp:2133
ReLU6 activation function.
Definition activation.hpp:2046
ReLU activation function.
Definition activation.hpp:1963
SELU activation function.
Definition activation.hpp:2233
Sigmoid Linear Unit activation function.
Definition activation.hpp:2368
Sigmoid activation function.
Definition activation.hpp:2316
Softmax activation function.
Definition activation.hpp:2422
Softmin activation function.
Definition activation.hpp:2507
Softplus activation function.
Definition activation.hpp:2592
Softshrink activation function.
Definition activation.hpp:2689
Softsign activation function.
Definition activation.hpp:2775
Tanh activation function.
Definition activation.hpp:2827
Tanhshrink activation function.
Definition activation.hpp:2879
Threshold activation function.
Definition activation.hpp:2935
Core components.
auto zip(T &&...seqs)
Provides the zip operation.
Definition zip.hpp:135
Definition core.hpp:73
nn_init
Enumerator for specifying the initialization of network weights.
Definition generator.hpp:30
activation
Enumerator for nonlinear activation functions.
Definition activation.hpp:26
short int short_t
Signed short integer type used by IgANet's compact enumerations.
Definition core.hpp:76
STL namespace.
Options.
Zip utility function.