@@ -237,10 +237,13 @@ def forward(
237237 query = self ._split_heads (self .W_query (inputs ))
238238 projected_key = self ._split_heads (self .W_key (key ))
239239 projected_value = self ._split_heads (self .W_value (value ))
240- scores = torch .matmul (
241- query ,
242- torch .tanh (projected_key .transpose (- 2 , - 1 )),
243- ) / self .score_scale
240+ scores = (
241+ torch .matmul (
242+ query ,
243+ torch .tanh (projected_key .transpose (- 2 , - 1 )),
244+ )
245+ / self .score_scale
246+ )
244247 weights = torch .softmax (scores , dim = - 1 )
245248 weights = torch .where (
246249 query_mask .transpose (1 , 2 ).unsqueeze (0 ) != 0 ,
@@ -316,9 +319,7 @@ def __init__(
316319 max_positions : int ,
317320 ) -> None :
318321 super ().__init__ ()
319- self .proj_in = Conv1dProjection (
320- latent_channels , hidden_channels , 1 , bias = False
321- )
322+ self .proj_in = Conv1dProjection (latent_channels , hidden_channels , 1 , bias = False )
322323 self .time_encoder = TimeEncoder (time_dim , time_hidden_channels )
323324 blocks : list [nn .Module ] = []
324325 for block_index in range (num_main_blocks ):
@@ -384,13 +385,9 @@ def forward(
384385 for block_index in range (self .num_main_blocks ):
385386 offset = block_index * 6
386387 hidden = self .main_blocks [offset ](hidden , latent_mask )
387- hidden = self .main_blocks [offset + 1 ](
388- hidden , time_embedding , latent_mask
389- )
388+ hidden = self .main_blocks [offset + 1 ](hidden , time_embedding , latent_mask )
390389 hidden = self .main_blocks [offset + 2 ](hidden , latent_mask )
391- hidden = self .main_blocks [offset + 3 ](
392- hidden , text , latent_mask , text_mask
393- )
390+ hidden = self .main_blocks [offset + 3 ](hidden , text , latent_mask , text_mask )
394391 hidden = self .main_blocks [offset + 4 ](hidden , latent_mask )
395392 hidden = self .main_blocks [offset + 5 ](
396393 hidden , style_key , style_value , latent_mask
@@ -422,15 +419,11 @@ def __init__(
422419 super ().__init__ ()
423420 if config .ttl .latent_dim <= 0 or config .ttl .chunk_compress_factor <= 0 :
424421 raise ValueError ("config.ttl dimensions must be positive" )
425- latent_channels = (
426- config .ttl .latent_dim * config .ttl .chunk_compress_factor
427- )
422+ latent_channels = config .ttl .latent_dim * config .ttl .chunk_compress_factor
428423 self .uncond_masker = UnconditionalMasker (
429424 text_channels , style_tokens , style_channels
430425 )
431- self .style_key = nn .Parameter (
432- torch .randn (1 , style_tokens , style_channels )
433- )
426+ self .style_key = nn .Parameter (torch .randn (1 , style_tokens , style_channels ))
434427 self .vector_field = VectorField (
435428 latent_channels ,
436429 hidden_channels ,
@@ -453,7 +446,7 @@ def __init__(
453446 self .style_channels = style_channels
454447 self .max_positions = max_positions
455448
456- def _validate_inputs (
449+ def _validate_inputs ( # noqa: C901
457450 self ,
458451 noisy_latent : torch .Tensor ,
459452 text_emb : torch .Tensor ,
@@ -463,17 +456,12 @@ def _validate_inputs(
463456 current_step : torch .Tensor ,
464457 total_step : torch .Tensor ,
465458 ) -> None :
466- if (
467- noisy_latent .ndim != 3
468- or noisy_latent .shape [1 ] != self .latent_channels
469- ):
459+ if noisy_latent .ndim != 3 or noisy_latent .shape [1 ] != self .latent_channels :
470460 raise ValueError (
471461 f"noisy_latent must have shape [B, { self .latent_channels } , L]"
472462 )
473463 if text_emb .ndim != 3 or text_emb .shape [1 ] != self .text_channels :
474- raise ValueError (
475- f"text_emb must have shape [B, { self .text_channels } , T]"
476- )
464+ raise ValueError (f"text_emb must have shape [B, { self .text_channels } , T]" )
477465 if style_ttl .ndim != 3 or style_ttl .shape [1 :] != (
478466 self .style_tokens ,
479467 self .style_channels ,
@@ -521,9 +509,7 @@ def _validate_inputs(
521509 if not torch .all (
522510 torch .isfinite (latent_valid_counts ) & (latent_valid_counts > 0 )
523511 ).item ():
524- raise ValueError (
525- "latent_mask must contain a valid position per sample"
526- )
512+ raise ValueError ("latent_mask must contain a valid position per sample" )
527513 text_valid_counts = text_mask .sum (dim = (1 , 2 ))
528514 if not torch .all (
529515 torch .isfinite (text_valid_counts ) & (text_valid_counts > 0 )
@@ -565,18 +551,14 @@ def forward(
565551 style_key = torch .cat (
566552 (
567553 self .style_key .expand (batch , - 1 , - 1 ),
568- self .uncond_masker .style_key_special_token .expand (
569- batch , - 1 , - 1
570- ),
554+ self .uncond_masker .style_key_special_token .expand (batch , - 1 , - 1 ),
571555 ),
572556 dim = 0 ,
573557 )
574558 style_value = torch .cat (
575559 (
576560 style_ttl ,
577- self .uncond_masker .style_value_special_token .expand (
578- batch , - 1 , - 1
579- ),
561+ self .uncond_masker .style_value_special_token .expand (batch , - 1 , - 1 ),
580562 ),
581563 dim = 0 ,
582564 )
0 commit comments