GPUCodeForces/S1/uucoco_#52/RationalFunctionApproximato...

34 lines
802 B
Python

import torch
import torch.nn as nn
class Model(nn.Module):
def __init__(self):
super().__init__()
self.a0 = nn.Parameter(torch.tensor(0.0))
self.a1 = nn.Parameter(torch.tensor(1.0))
self.a2 = nn.Parameter(torch.tensor(0.0))
self.b1 = nn.Parameter(torch.tensor(0.0))
self.b2 = nn.Parameter(torch.tensor(0.0))
def forward(self, x: torch.Tensor) -> torch.Tensor:
x_sq = x * x
abs_x = torch.abs(x)
num = self.a0 + self.a1 * x + self.a2 * x_sq
den = 1.0 + torch.abs(self.b1) * abs_x + torch.abs(self.b2) * x_sq
return num / den
batch_size = 1024
feature_dim = 4096
def get_inputs():
x = torch.randn(batch_size, feature_dim, dtype=torch.float32)
return [x]
def get_init_inputs():
return []