Я знаю, что JAX вполне функционален, и, вероятно, я использовал объектно-ориентированный стиль в PyTorch, когда спрашивал об этом, но я просто не знаю. Я не знаю, как это сделать в JAX.
Итак, в PyTorch я обычно определял некоторые пользовательские блоки следующим образом:
Предположим, у меня есть такой супер-модный пользовательский слой, как этот. :
Код: Выделить всё
def net_block(self,x):
U,S,V = torch.svd(x)
return torch.sin(torch.relu(S))
Код: Выделить всё
class Network(nn.Module):
def __init__(self, x):
super(Network, self).__init__()
self.model = nn.Sequential(
net_block(),
net_block(),
net_block()
)
def forward(self, X):
return self.model(X)
Но в этом случае В этом случае мой код будет совершенно нечитабельным в JAX.
Я слышал о Flax, поэтому думаю, что net_block может быть реализован во Flax, но вопрос в том, как я могу каскадировать net_block< /code> в мою основную модель (например, как я это сделал в Python?)
А потом, как я могу вычислять градиенты?
Я слышал о Flax, поэтому я думаю, что net_block может быть реализован с помощью Flax, но вопрос в том, как я могу каскадировать net_block в свою основную модель (например, как я это сделал в Python?)
И как тогда я могу вычислять градиенты?
Во-вторых, как я могу обрабатывать пакетную обработку в этом случае для JAX? Позаботится ли Флакс об этом сам?
Подробнее здесь: https://stackoverflow.com/questions/790 ... ent-in-jax