Belle II Software light-2609-luna
MultiClassNet Class Reference
Inheritance diagram for MultiClassNet:
Collaboration diagram for MultiClassNet:

Public Member Functions

 __init__ (self, input_size, num_labels=3)
 
 forward (self, x)
 

Public Attributes

 activation = nn.ReLU()
 Activation function applied between all linear layers.
 
 network
 Fully connected network layers.
 

Protected Member Functions

 _init_weights (self)
 

Detailed Description

Multi-class classification network with fully connected layers.

Two architectures selected automatically by num_labels:
- Deep (num_labels <= 10, e.g. category network with 3 classes):
  256 -> 128 (dropout 0.2) -> 64 (dropout 0.15) -> 32 (dropout 0.1) -> 16 -> num_labels
- Shallow (num_labels > 10, e.g. main network with 139 classes):
  256 -> 128 (dropout 0.1) -> 128 (dropout 0.1) -> 64 (dropout 0.05) -> num_labels
- ReLU activation, Xavier initialization (gain=0.5, bias=0.01)

Definition at line 710 of file train.py.

Constructor & Destructor Documentation

◆ __init__()

__init__ ( self,
input_size,
num_labels = 3 )
Build the network graph for the given input size and number of output classes.

Definition at line 721 of file train.py.

721 def __init__(self, input_size, num_labels=3):
722 """Build the network graph for the given input size and number of output classes."""
723 super().__init__()
724
725
726 self.activation = nn.ReLU()
727
728 if num_labels <= 10:
729
730 self.network = nn.Sequential(
731 nn.Linear(input_size, 128),
732 self.activation,
733 nn.Linear(128, 128),
734 nn.Dropout(0.2),
735 self.activation,
736 nn.Linear(128, 64),
737 nn.Dropout(0.15),
738 self.activation,
739 nn.Linear(64, 32),
740 nn.Dropout(0.1),
741 self.activation,
742 nn.Linear(32, 16),
743 self.activation,
744 nn.Linear(16, num_labels)
745 )
746 else:
747 self.network = nn.Sequential(
748 nn.Linear(input_size, 256),
749 self.activation,
750 nn.Linear(256, 128),
751 nn.Dropout(0.1),
752 self.activation,
753 nn.Linear(128, 128),
754 nn.Dropout(0.1),
755 self.activation,
756 nn.Linear(128, 64),
757 nn.Dropout(0.05),
758 self.activation,
759 nn.Linear(64, num_labels)
760 )
761
762 self._init_weights()
763

Member Function Documentation

◆ _init_weights()

_init_weights ( self)
protected
Initialize weights with Xavier uniform (gain=0.5) and bias=0.01.

Definition at line 764 of file train.py.

764 def _init_weights(self):
765 """Initialize weights with Xavier uniform (gain=0.5) and bias=0.01."""
766 for m in self.modules():
767 if isinstance(m, nn.Linear):
768 nn.init.xavier_uniform_(m.weight, gain=0.5)
769 nn.init.constant_(m.bias, 0.01)
770

◆ forward()

forward ( self,
x )
Run a forward pass through the network.

Definition at line 771 of file train.py.

771 def forward(self, x):
772 """Run a forward pass through the network."""
773 return self.network(x)
774
775

Member Data Documentation

◆ activation

activation = nn.ReLU()

Activation function applied between all linear layers.

Definition at line 726 of file train.py.

◆ network

network
Initial value:
= nn.Sequential(
nn.Linear(input_size, 128),
self.activation,
nn.Linear(128, 128),
nn.Dropout(0.2),
self.activation,
nn.Linear(128, 64),
nn.Dropout(0.15),
self.activation,
nn.Linear(64, 32),
nn.Dropout(0.1),
self.activation,
nn.Linear(32, 16),
self.activation,
nn.Linear(16, num_labels)
)

Fully connected network layers.

Definition at line 730 of file train.py.


The documentation for this class was generated from the following file: