| 60 | |
| 61 | # helper for batch norm |
| 62 | def batch_norm( |
| 63 | input_, |
| 64 | gamma, |
| 65 | beta, |
| 66 | running_mean, |
| 67 | running_var, |
| 68 | is_training, |
| 69 | axes='per-activation'): |
| 70 | |
| 71 | if is_training: |
| 72 | # returns: |
| 73 | # batch-normalized output |
| 74 | # batch mean |
| 75 | # batch variance |
| 76 | # running mean (for later use as population mean estimate) |
| 77 | # running var (for later use as population var estimate) |
| 78 | out, _, _, new_running_mean, new_running_var = batch_normalization_train( |
| 79 | input_, |
| 80 | gamma, |
| 81 | beta, |
| 82 | running_mean=running_mean, |
| 83 | running_var=running_var, |
| 84 | axes=axes, |
| 85 | running_average_factor=0.9, |
| 86 | ) |
| 87 | else: |
| 88 | new_running_mean = None |
| 89 | new_running_var = None # just to ensure we don't try to use them |
| 90 | out = batch_normalization_test( |
| 91 | input_, |
| 92 | gamma, |
| 93 | beta, |
| 94 | running_mean, |
| 95 | running_var, |
| 96 | axes=axes, |
| 97 | ) |
| 98 | return out, new_running_mean, new_running_var |
| 99 | |
| 100 | |
| 101 | class ConvLayer: |