File size: 1,893 Bytes
59630ba | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 | from abc import abstractmethod, ABC
from typing import Optional
import torch
from torch import nn
from omegaconf import DictConfig
from .modules.embeddings import (
StochasticTimeEmbedding,
RandomDropoutCondEmbedding,
)
class BaseBackbone(ABC, nn.Module):
def __init__(
self,
cfg: DictConfig,
x_shape: torch.Size,
max_tokens: int,
external_cond_dim: int,
use_causal_mask=True,
):
super().__init__()
self.cfg = cfg
self.external_cond_dim = external_cond_dim
self.use_causal_mask = use_causal_mask
self.x_shape = x_shape
self.noise_level_pos_embedding = StochasticTimeEmbedding(
dim=self.noise_level_dim,
time_embed_dim=self.noise_level_emb_dim,
use_fourier=self.cfg.get("use_fourier_noise_embedding", False),
)
self.external_cond_embedding = self._build_external_cond_embedding()
def _build_external_cond_embedding(self) -> Optional[nn.Module]:
return (
RandomDropoutCondEmbedding(
self.external_cond_dim,
self.external_cond_emb_dim,
dropout_prob=self.cfg.get("external_cond_dropout", 0.0),
)
if self.external_cond_dim
else None
)
@property
def noise_level_dim(self):
return max(self.noise_level_emb_dim // 4, 32)
@property
@abstractmethod
def noise_level_emb_dim(self):
raise NotImplementedError
@property
@abstractmethod
def external_cond_emb_dim(self):
raise NotImplementedError
@abstractmethod
def forward(
self,
x: torch.Tensor,
noise_levels: torch.Tensor,
external_cond: Optional[torch.Tensor] = None,
external_cond_mask: Optional[torch.Tensor] = None,
):
raise NotImplementedError
|