yield-estimation / modeling_yield.py
Sarikaa-Sridhar's picture
Rename modeling_yield_text.py to modeling_yield.py
9fff2c7 verified
Raw
History Blame Contribute Delete
5.26 kB
import torch
from torch import nn
from transformers import PreTrainedModel
from transformers.modeling_outputs import SequenceClassifierOutput
from .configuration_yield import YieldConfig
from .yield_transformer import YieldTransformer
class YieldForSequenceClassification(PreTrainedModel):
config_class = YieldConfig
base_model_prefix = "yield_model"
def __init__(self, config: YieldConfig):
super().__init__(config)
self.yield_model = YieldTransformer(
w_dim=config.W,
soil_dim=config.S,
d_model=config.d_model,
nhead=config.nhead,
num_layers=config.num_layers,
dim_ff=config.dim_ff,
dropout=config.dropout,
use_crop=config.use_crop,
crop_emb_dim=config.crop_emb_dim,
max_weeks=max(52, config.K),
pool=config.pool,
)
self.post_init()
def forward(
self,
weather,
soil,
crop_id,
horizon_idx=None,
labels=None,
**kwargs,
):
# ==================================================
# Shape checks
# ==================================================
if weather.ndim != 3:
raise ValueError(
f"weather must have shape [B,K,W], "
f"received {tuple(weather.shape)}"
)
if soil.ndim != 2:
raise ValueError(
f"soil must have shape [B,S], "
f"received {tuple(soil.shape)}"
)
# ==================================================
# Training normalization statistics
# ==================================================
w_mean = torch.tensor(
self.config.w_mean,
device=weather.device,
dtype=weather.dtype,
)
w_std = torch.tensor(
self.config.w_std,
device=weather.device,
dtype=weather.dtype,
)
s_mean = torch.tensor(
self.config.s_mean,
device=soil.device,
dtype=soil.dtype,
)
s_std = torch.tensor(
self.config.s_std,
device=soil.device,
dtype=soil.dtype,
)
# NaN -> training mean
weather = torch.where(
torch.isnan(weather),
w_mean.view(1, 1, -1),
weather,
)
soil = torch.where(
torch.isnan(soil),
s_mean.view(1, -1),
soil,
)
# normalize
weather = (
weather - w_mean.view(1, 1, -1)
) / w_std.view(1, 1, -1)
soil = (
soil - s_mean.view(1, -1)
) / s_std.view(1, -1)
# ==================================================
# Cutoff
# ==================================================
if horizon_idx is None:
horizon_idx = torch.full(
(weather.shape[0],),
weather.shape[1],
dtype=torch.long,
device=weather.device,
)
if not torch.is_tensor(horizon_idx):
horizon_idx = torch.tensor(
horizon_idx,
dtype=torch.long,
device=weather.device,
)
horizon_idx = horizon_idx.to(
weather.device
).long()
if horizon_idx.ndim == 0:
horizon_idx = horizon_idx.unsqueeze(0)
# FlexServe requests are normally one sample.
# Ensure one common temporal length for a batch.
unique_cutoffs = torch.unique(
horizon_idx
)
if len(unique_cutoffs) != 1:
raise ValueError(
"All samples in one batch must use the same cutoff."
)
t_eff = int(
unique_cutoffs[0].item()
)
weather = weather[
:,
:t_eff,
:
]
# ==================================================
# Existing trained model
# ==================================================
logits_norm = self.yield_model(
weather,
soil,
crop_id,
horizon_idx=horizon_idx,
causal=True,
return_sequence=False,
)
# Restore yield to bu/acre.
y_mean = torch.tensor(
self.config.y_mean,
device=logits_norm.device,
dtype=logits_norm.dtype,
)
y_std = torch.tensor(
self.config.y_std,
device=logits_norm.device,
dtype=logits_norm.dtype,
)
predicted_yield = (
logits_norm * y_std
+ y_mean
)
# TextClassificationPipeline expects [B, num_labels].
logits = predicted_yield.unsqueeze(-1)
loss = None
if labels is not None:
labels = labels.to(
logits.dtype
).view(-1)
loss = nn.functional.mse_loss(
predicted_yield,
labels,
)
return SequenceClassifierOutput(
loss=loss,
logits=logits,
)