ARotting's picture
Publish First-order MAML, pooled, and random sine initializations
47dfc70 verified
Raw
History Blame Contribute Delete
566 Bytes
from __future__ import annotations
import torch
from torch import nn
class SineRegressor(nn.Module):
def __init__(self) -> None:
super().__init__()
self.network = nn.Sequential(
nn.Linear(1, 40),
nn.Tanh(),
nn.Linear(40, 40),
nn.Tanh(),
nn.Linear(40, 1),
)
def forward(self, inputs: torch.Tensor) -> torch.Tensor:
return self.network(inputs)
def parameter_count(module: nn.Module) -> int:
return sum(parameter.numel() for parameter in module.parameters())