(embedding_type, embedding_dim, embedding_scale=10000)
| 76 | |
| 77 | |
| 78 | def get_timestep_embedding(embedding_type, embedding_dim, embedding_scale=10000): |
| 79 | if embedding_type == 'sinusoidal': |
| 80 | emb_func = SinusoidalEmbedding(embedding_dim, scale=embedding_scale) |
| 81 | elif embedding_type == 'fourier': |
| 82 | emb_func = GaussianFourierEmbedding(embedding_size=embedding_dim, scale=embedding_scale) |
| 83 | else: |
| 84 | raise NotImplementedError |
| 85 | return emb_func |
no test coverage detected