MCPcopy Create free account
hub / github.com/openai/point-e / CLIPImagePointDiffusionTransformer

Class CLIPImagePointDiffusionTransformer

point_e/models/transformer.py:229–287  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

227
228
229class CLIPImagePointDiffusionTransformer(PointDiffusionTransformer):
230 def __init__(
231 self,
232 *,
233 device: torch.device,
234 dtype: torch.dtype,
235 n_ctx: int = 1024,
236 token_cond: bool = False,
237 cond_drop_prob: float = 0.0,
238 frozen_clip: bool = True,
239 cache_dir: Optional[str] = None,
240 **kwargs,
241 ):
242 super().__init__(device=device, dtype=dtype, n_ctx=n_ctx + int(token_cond), **kwargs)
243 self.n_ctx = n_ctx
244 self.token_cond = token_cond
245 self.clip = (FrozenImageCLIP if frozen_clip else ImageCLIP)(device, cache_dir=cache_dir)
246 self.clip_embed = nn.Linear(
247 self.clip.feature_dim, self.backbone.width, device=device, dtype=dtype
248 )
249 self.cond_drop_prob = cond_drop_prob
250
251 def cached_model_kwargs(self, batch_size: int, model_kwargs: Dict[str, Any]) -> Dict[str, Any]:
252 with torch.no_grad():
253 return dict(embeddings=self.clip(batch_size, **model_kwargs))
254
255 def forward(
256 self,
257 x: torch.Tensor,
258 t: torch.Tensor,
259 images: Optional[Iterable[Optional[ImageType]]] = None,
260 texts: Optional[Iterable[Optional[str]]] = None,
261 embeddings: Optional[Iterable[Optional[torch.Tensor]]] = None,
262 ):
263 """
264 :param x: an [N x C x T] tensor.
265 :param t: an [N] tensor.
266 :param images: a batch of images to condition on.
267 :param texts: a batch of texts to condition on.
268 :param embeddings: a batch of CLIP embeddings to condition on.
269 :return: an [N x C' x T] tensor.
270 """
271 assert x.shape[-1] == self.n_ctx
272
273 t_embed = self.time_embed(timestep_embedding(t, self.backbone.width))
274 clip_out = self.clip(batch_size=len(x), images=images, texts=texts, embeddings=embeddings)
275 assert len(clip_out.shape) == 2 and clip_out.shape[0] == x.shape[0]
276
277 if self.training:
278 mask = torch.rand(size=[len(x)]) >= self.cond_drop_prob
279 clip_out = clip_out * mask[:, None].to(clip_out)
280
281 # Rescale the features to have unit variance
282 clip_out = math.sqrt(clip_out.shape[1]) * clip_out
283
284 clip_embed = self.clip_embed(clip_out)
285
286 cond = [(clip_embed, self.token_cond), (t_embed, self.time_token_cond)]

Callers 1

model_from_configFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected