(self, x, timesteps=None, context=None)
| 324 | ] |
| 325 | |
| 326 | def __call__(self, x, timesteps=None, context=None): |
| 327 | # TODO: real time embedding |
| 328 | t_emb = timestep_embedding(timesteps, 320) |
| 329 | emb = t_emb.sequential(self.time_embed) |
| 330 | |
| 331 | |
| 332 | |
| 333 | def run(x, bb): |
| 334 | if isinstance(bb, ResBlock): x = bb(x, emb) |
| 335 | elif isinstance(bb, SpatialTransformer): x = bb(x, context) |
| 336 | else: x = bb(x) |
| 337 | return x |
| 338 | |
| 339 | saved_inputs = [] |
| 340 | for i,b in enumerate(self.input_blocks): |
| 341 | for bb in b: |
| 342 | x = run(x, bb) |
| 343 | saved_inputs.append(x) |
| 344 | for bb in self.middle_block: |
| 345 | x = run(x, bb) |
| 346 | for i,b in enumerate(self.output_blocks): |
| 347 | x = x.cat(saved_inputs.pop(), dim=1) |
| 348 | for bb in b: |
| 349 | x = run(x, bb) |
| 350 | return x.sequential(self.out) |
| 351 | |
| 352 | class CLIPMLP: |
| 353 | def __init__(self): |
nothing calls this directly
no test coverage detected