神经网络类中激活函数参数选择的简洁实现方法咨询
Great question! Those endless if/else checks for activation functions are such a code smell—they make the code harder to read and a pain to maintain when you add new functions. Here are a couple of clean, Pythonic solutions to fix this:
1. Use a Function Mapping Dictionary
If you still want to let users specify the activation function via a string (like your current 'sig' default), create a dictionary that maps string identifiers to your class's activation methods. This replaces all those conditional checks with a single dictionary lookup:
import numpy as np class NeuralNetworks: def __init__(self, func='sig'): # Define a mapping of string keys to activation methods self._activation_map = { 'sig': self.sigmoid, 'relu': self.reLu, # Add new functions here as you implement them } # Validate the input to avoid typos if func not in self._activation_map: valid_funcs = ', '.join(self._activation_map.keys()) raise ValueError(f"Invalid activation function. Choose from: {valid_funcs}") # Assign the selected function directly to self.func self.func = self._activation_map[func] def sigmoid(self, x): return 1 / (1 + np.exp(-x)) def reLu(self, x): return np.maximum(0, x) def propagate(self, x): # No more if/else! Just call self.func directly activated_output = self.func(x) # Rest of your propagation logic...
This approach keeps your propagate (and other methods) clean. Adding a new activation function only requires updating the _activation_map—no need to touch every method that uses the function.
2. Pass Function Objects Directly (Recommended)
For even more flexibility, let users pass the activation function itself instead of a string. This eliminates the mapping entirely and lets users use custom functions outside your class:
import numpy as np class NeuralNetworks: def __init__(self, func=None): # Default to sigmoid if no function is provided self.func = func if func is not None else self.sigmoid def sigmoid(self, x): return 1 / (1 + np.exp(-x)) def reLu(self, x): return np.maximum(0, x) def propagate(self, x): activated_output = self.func(x) # Rest of your propagation logic... # Usage examples: # Use the built-in ReLU nn_relu = NeuralNetworks(func=NeuralNetworks.reLu) # Or pass a custom activation function def custom_tanh(x): return np.tanh(x) nn_custom = NeuralNetworks(func=custom_tanh)
This is my favorite approach because it's simpler and more flexible. Users aren't limited to the functions you've pre-defined in the class, and your code stays lean with zero conditional checks for activation functions.
If you ever need to serialize or log which activation function is being used (e.g., for saving the model), you can add a small attribute to track the function's name—for example, storing self.func_name alongside self.func in either approach.
内容的提问来源于stack exchange,提问作者David Montgomery

