CoolFace
Apppublic

declare-lab/tango2

sourceHugging Faceupdated 2y agoView on Hugging Face
92likes
unet_1d_blocks.py669 linesDownload Raw Back to models
1# Copyright 2023 The HuggingFace Team. All rights reserved.2#3# Licensed under the Apache License, Version 2.0 (the "License");4# you may not use this file except in compliance with the License.5# You may obtain a copy of the License at6#7#     http://www.apache.org/licenses/LICENSE-2.08#9# Unless required by applicable law or agreed to in writing, software10# distributed under the License is distributed on an "AS IS" BASIS,11# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.12# See the License for the specific language governing permissions and13# limitations under the License.14import math15 16import torch17import torch.nn.functional as F18from torch import nn19 20from .resnet import Downsample1D, ResidualTemporalBlock1D, Upsample1D, rearrange_dims21 22 23class DownResnetBlock1D(nn.Module):24    def __init__(25        self,26        in_channels,27        out_channels=None,28        num_layers=1,29        conv_shortcut=False,30        temb_channels=32,31        groups=32,32        groups_out=None,33        non_linearity=None,34        time_embedding_norm="default",35        output_scale_factor=1.0,36        add_downsample=True,37    ):38        super().__init__()39        self.in_channels = in_channels40        out_channels = in_channels if out_channels is None else out_channels41        self.out_channels = out_channels42        self.use_conv_shortcut = conv_shortcut43        self.time_embedding_norm = time_embedding_norm44        self.add_downsample = add_downsample45        self.output_scale_factor = output_scale_factor46 47        if groups_out is None:48            groups_out = groups49 50        # there will always be at least one resnet51        resnets = [ResidualTemporalBlock1D(in_channels, out_channels, embed_dim=temb_channels)]52 53        for _ in range(num_layers):54            resnets.append(ResidualTemporalBlock1D(out_channels, out_channels, embed_dim=temb_channels))55 56        self.resnets = nn.ModuleList(resnets)57 58        if non_linearity == "swish":59            self.nonlinearity = lambda x: F.silu(x)60        elif non_linearity == "mish":61            self.nonlinearity = nn.Mish()62        elif non_linearity == "silu":63            self.nonlinearity = nn.SiLU()64        else:65            self.nonlinearity = None66 67        self.downsample = None68        if add_downsample:69            self.downsample = Downsample1D(out_channels, use_conv=True, padding=1)70 71    def forward(self, hidden_states, temb=None):72        output_states = ()73 74        hidden_states = self.resnets[0](hidden_states, temb)75        for resnet in self.resnets[1:]:76            hidden_states = resnet(hidden_states, temb)77 78        output_states += (hidden_states,)79 80        if self.nonlinearity is not None:81            hidden_states = self.nonlinearity(hidden_states)82 83        if self.downsample is not None:84            hidden_states = self.downsample(hidden_states)85 86        return hidden_states, output_states87 88 89class UpResnetBlock1D(nn.Module):90    def __init__(91        self,92        in_channels,93        out_channels=None,94        num_layers=1,95        temb_channels=32,96        groups=32,97        groups_out=None,98        non_linearity=None,99        time_embedding_norm="default",100        output_scale_factor=1.0,101        add_upsample=True,102    ):103        super().__init__()104        self.in_channels = in_channels105        out_channels = in_channels if out_channels is None else out_channels106        self.out_channels = out_channels107        self.time_embedding_norm = time_embedding_norm108        self.add_upsample = add_upsample109        self.output_scale_factor = output_scale_factor110 111        if groups_out is None:112            groups_out = groups113 114        # there will always be at least one resnet115        resnets = [ResidualTemporalBlock1D(2 * in_channels, out_channels, embed_dim=temb_channels)]116 117        for _ in range(num_layers):118            resnets.append(ResidualTemporalBlock1D(out_channels, out_channels, embed_dim=temb_channels))119 120        self.resnets = nn.ModuleList(resnets)121 122        if non_linearity == "swish":123            self.nonlinearity = lambda x: F.silu(x)124        elif non_linearity == "mish":125            self.nonlinearity = nn.Mish()126        elif non_linearity == "silu":127            self.nonlinearity = nn.SiLU()128        else:129            self.nonlinearity = None130 131        self.upsample = None132        if add_upsample:133            self.upsample = Upsample1D(out_channels, use_conv_transpose=True)134 135    def forward(self, hidden_states, res_hidden_states_tuple=None, temb=None):136        if res_hidden_states_tuple is not None:137            res_hidden_states = res_hidden_states_tuple[-1]138            hidden_states = torch.cat((hidden_states, res_hidden_states), dim=1)139 140        hidden_states = self.resnets[0](hidden_states, temb)141        for resnet in self.resnets[1:]:142            hidden_states = resnet(hidden_states, temb)143 144        if self.nonlinearity is not None:145            hidden_states = self.nonlinearity(hidden_states)146 147        if self.upsample is not None:148            hidden_states = self.upsample(hidden_states)149 150        return hidden_states151 152 153class ValueFunctionMidBlock1D(nn.Module):154    def __init__(self, in_channels, out_channels, embed_dim):155        super().__init__()156        self.in_channels = in_channels157        self.out_channels = out_channels158        self.embed_dim = embed_dim159 160        self.res1 = ResidualTemporalBlock1D(in_channels, in_channels // 2, embed_dim=embed_dim)161        self.down1 = Downsample1D(out_channels // 2, use_conv=True)162        self.res2 = ResidualTemporalBlock1D(in_channels // 2, in_channels // 4, embed_dim=embed_dim)163        self.down2 = Downsample1D(out_channels // 4, use_conv=True)164 165    def forward(self, x, temb=None):166        x = self.res1(x, temb)167        x = self.down1(x)168        x = self.res2(x, temb)169        x = self.down2(x)170        return x171 172 173class MidResTemporalBlock1D(nn.Module):174    def __init__(175        self,176        in_channels,177        out_channels,178        embed_dim,179        num_layers: int = 1,180        add_downsample: bool = False,181        add_upsample: bool = False,182        non_linearity=None,183    ):184        super().__init__()185        self.in_channels = in_channels186        self.out_channels = out_channels187        self.add_downsample = add_downsample188 189        # there will always be at least one resnet190        resnets = [ResidualTemporalBlock1D(in_channels, out_channels, embed_dim=embed_dim)]191 192        for _ in range(num_layers):193            resnets.append(ResidualTemporalBlock1D(out_channels, out_channels, embed_dim=embed_dim))194 195        self.resnets = nn.ModuleList(resnets)196 197        if non_linearity == "swish":198            self.nonlinearity = lambda x: F.silu(x)199        elif non_linearity == "mish":200            self.nonlinearity = nn.Mish()201        elif non_linearity == "silu":202            self.nonlinearity = nn.SiLU()203        else:204            self.nonlinearity = None205 206        self.upsample = None207        if add_upsample:208            self.upsample = Downsample1D(out_channels, use_conv=True)209 210        self.downsample = None211        if add_downsample:212            self.downsample = Downsample1D(out_channels, use_conv=True)213 214        if self.upsample and self.downsample:215            raise ValueError("Block cannot downsample and upsample")216 217    def forward(self, hidden_states, temb):218        hidden_states = self.resnets[0](hidden_states, temb)219        for resnet in self.resnets[1:]:220            hidden_states = resnet(hidden_states, temb)221 222        if self.upsample:223            hidden_states = self.upsample(hidden_states)224        if self.downsample:225            self.downsample = self.downsample(hidden_states)226 227        return hidden_states228 229 230class OutConv1DBlock(nn.Module):231    def __init__(self, num_groups_out, out_channels, embed_dim, act_fn):232        super().__init__()233        self.final_conv1d_1 = nn.Conv1d(embed_dim, embed_dim, 5, padding=2)234        self.final_conv1d_gn = nn.GroupNorm(num_groups_out, embed_dim)235        if act_fn == "silu":236            self.final_conv1d_act = nn.SiLU()237        if act_fn == "mish":238            self.final_conv1d_act = nn.Mish()239        self.final_conv1d_2 = nn.Conv1d(embed_dim, out_channels, 1)240 241    def forward(self, hidden_states, temb=None):242        hidden_states = self.final_conv1d_1(hidden_states)243        hidden_states = rearrange_dims(hidden_states)244        hidden_states = self.final_conv1d_gn(hidden_states)245        hidden_states = rearrange_dims(hidden_states)246        hidden_states = self.final_conv1d_act(hidden_states)247        hidden_states = self.final_conv1d_2(hidden_states)248        return hidden_states249 250 251class OutValueFunctionBlock(nn.Module):252    def __init__(self, fc_dim, embed_dim):253        super().__init__()254        self.final_block = nn.ModuleList(255            [256                nn.Linear(fc_dim + embed_dim, fc_dim // 2),257                nn.Mish(),258                nn.Linear(fc_dim // 2, 1),259            ]260        )261 262    def forward(self, hidden_states, temb):263        hidden_states = hidden_states.view(hidden_states.shape[0], -1)264        hidden_states = torch.cat((hidden_states, temb), dim=-1)265        for layer in self.final_block:266            hidden_states = layer(hidden_states)267 268        return hidden_states269 270 271_kernels = {272    "linear": [1 / 8, 3 / 8, 3 / 8, 1 / 8],273    "cubic": [-0.01171875, -0.03515625, 0.11328125, 0.43359375, 0.43359375, 0.11328125, -0.03515625, -0.01171875],274    "lanczos3": [275        0.003689131001010537,276        0.015056144446134567,277        -0.03399861603975296,278        -0.066637322306633,279        0.13550527393817902,280        0.44638532400131226,281        0.44638532400131226,282        0.13550527393817902,283        -0.066637322306633,284        -0.03399861603975296,285        0.015056144446134567,286        0.003689131001010537,287    ],288}289 290 291class Downsample1d(nn.Module):292    def __init__(self, kernel="linear", pad_mode="reflect"):293        super().__init__()294        self.pad_mode = pad_mode295        kernel_1d = torch.tensor(_kernels[kernel])296        self.pad = kernel_1d.shape[0] // 2 - 1297        self.register_buffer("kernel", kernel_1d)298 299    def forward(self, hidden_states):300        hidden_states = F.pad(hidden_states, (self.pad,) * 2, self.pad_mode)301        weight = hidden_states.new_zeros([hidden_states.shape[1], hidden_states.shape[1], self.kernel.shape[0]])302        indices = torch.arange(hidden_states.shape[1], device=hidden_states.device)303        weight[indices, indices] = self.kernel.to(weight)304        return F.conv1d(hidden_states, weight, stride=2)305 306 307class Upsample1d(nn.Module):308    def __init__(self, kernel="linear", pad_mode="reflect"):309        super().__init__()310        self.pad_mode = pad_mode311        kernel_1d = torch.tensor(_kernels[kernel]) * 2312        self.pad = kernel_1d.shape[0] // 2 - 1313        self.register_buffer("kernel", kernel_1d)314 315    def forward(self, hidden_states, temb=None):316        hidden_states = F.pad(hidden_states, ((self.pad + 1) // 2,) * 2, self.pad_mode)317        weight = hidden_states.new_zeros([hidden_states.shape[1], hidden_states.shape[1], self.kernel.shape[0]])318        indices = torch.arange(hidden_states.shape[1], device=hidden_states.device)319        weight[indices, indices] = self.kernel.to(weight)320        return F.conv_transpose1d(hidden_states, weight, stride=2, padding=self.pad * 2 + 1)321 322 323class SelfAttention1d(nn.Module):324    def __init__(self, in_channels, n_head=1, dropout_rate=0.0):325        super().__init__()326        self.channels = in_channels327        self.group_norm = nn.GroupNorm(1, num_channels=in_channels)328        self.num_heads = n_head329 330        self.query = nn.Linear(self.channels, self.channels)331        self.key = nn.Linear(self.channels, self.channels)332        self.value = nn.Linear(self.channels, self.channels)333 334        self.proj_attn = nn.Linear(self.channels, self.channels, bias=True)335 336        self.dropout = nn.Dropout(dropout_rate, inplace=True)337 338    def transpose_for_scores(self, projection: torch.Tensor) -> torch.Tensor:339        new_projection_shape = projection.size()[:-1] + (self.num_heads, -1)340        # move heads to 2nd position (B, T, H * D) -> (B, T, H, D) -> (B, H, T, D)341        new_projection = projection.view(new_projection_shape).permute(0, 2, 1, 3)342        return new_projection343 344    def forward(self, hidden_states):345        residual = hidden_states346        batch, channel_dim, seq = hidden_states.shape347 348        hidden_states = self.group_norm(hidden_states)349        hidden_states = hidden_states.transpose(1, 2)350 351        query_proj = self.query(hidden_states)352        key_proj = self.key(hidden_states)353        value_proj = self.value(hidden_states)354 355        query_states = self.transpose_for_scores(query_proj)356        key_states = self.transpose_for_scores(key_proj)357        value_states = self.transpose_for_scores(value_proj)358 359        scale = 1 / math.sqrt(math.sqrt(key_states.shape[-1]))360 361        attention_scores = torch.matmul(query_states * scale, key_states.transpose(-1, -2) * scale)362        attention_probs = torch.softmax(attention_scores, dim=-1)363 364        # compute attention output365        hidden_states = torch.matmul(attention_probs, value_states)366 367        hidden_states = hidden_states.permute(0, 2, 1, 3).contiguous()368        new_hidden_states_shape = hidden_states.size()[:-2] + (self.channels,)369        hidden_states = hidden_states.view(new_hidden_states_shape)370 371        # compute next hidden_states372        hidden_states = self.proj_attn(hidden_states)373        hidden_states = hidden_states.transpose(1, 2)374        hidden_states = self.dropout(hidden_states)375 376        output = hidden_states + residual377 378        return output379 380 381class ResConvBlock(nn.Module):382    def __init__(self, in_channels, mid_channels, out_channels, is_last=False):383        super().__init__()384        self.is_last = is_last385        self.has_conv_skip = in_channels != out_channels386 387        if self.has_conv_skip:388            self.conv_skip = nn.Conv1d(in_channels, out_channels, 1, bias=False)389 390        self.conv_1 = nn.Conv1d(in_channels, mid_channels, 5, padding=2)391        self.group_norm_1 = nn.GroupNorm(1, mid_channels)392        self.gelu_1 = nn.GELU()393        self.conv_2 = nn.Conv1d(mid_channels, out_channels, 5, padding=2)394 395        if not self.is_last:396            self.group_norm_2 = nn.GroupNorm(1, out_channels)397            self.gelu_2 = nn.GELU()398 399    def forward(self, hidden_states):400        residual = self.conv_skip(hidden_states) if self.has_conv_skip else hidden_states401 402        hidden_states = self.conv_1(hidden_states)403        hidden_states = self.group_norm_1(hidden_states)404        hidden_states = self.gelu_1(hidden_states)405        hidden_states = self.conv_2(hidden_states)406 407        if not self.is_last:408            hidden_states = self.group_norm_2(hidden_states)409            hidden_states = self.gelu_2(hidden_states)410 411        output = hidden_states + residual412        return output413 414 415class UNetMidBlock1D(nn.Module):416    def __init__(self, mid_channels, in_channels, out_channels=None):417        super().__init__()418 419        out_channels = in_channels if out_channels is None else out_channels420 421        # there is always at least one resnet422        self.down = Downsample1d("cubic")423        resnets = [424            ResConvBlock(in_channels, mid_channels, mid_channels),425            ResConvBlock(mid_channels, mid_channels, mid_channels),426            ResConvBlock(mid_channels, mid_channels, mid_channels),427            ResConvBlock(mid_channels, mid_channels, mid_channels),428            ResConvBlock(mid_channels, mid_channels, mid_channels),429            ResConvBlock(mid_channels, mid_channels, out_channels),430        ]431        attentions = [432            SelfAttention1d(mid_channels, mid_channels // 32),433            SelfAttention1d(mid_channels, mid_channels // 32),434            SelfAttention1d(mid_channels, mid_channels // 32),435            SelfAttention1d(mid_channels, mid_channels // 32),436            SelfAttention1d(mid_channels, mid_channels // 32),437            SelfAttention1d(out_channels, out_channels // 32),438        ]439        self.up = Upsample1d(kernel="cubic")440 441        self.attentions = nn.ModuleList(attentions)442        self.resnets = nn.ModuleList(resnets)443 444    def forward(self, hidden_states, temb=None):445        hidden_states = self.down(hidden_states)446        for attn, resnet in zip(self.attentions, self.resnets):447            hidden_states = resnet(hidden_states)448            hidden_states = attn(hidden_states)449 450        hidden_states = self.up(hidden_states)451 452        return hidden_states453 454 455class AttnDownBlock1D(nn.Module):456    def __init__(self, out_channels, in_channels, mid_channels=None):457        super().__init__()458        mid_channels = out_channels if mid_channels is None else mid_channels459 460        self.down = Downsample1d("cubic")461        resnets = [462            ResConvBlock(in_channels, mid_channels, mid_channels),463            ResConvBlock(mid_channels, mid_channels, mid_channels),464            ResConvBlock(mid_channels, mid_channels, out_channels),465        ]466        attentions = [467            SelfAttention1d(mid_channels, mid_channels // 32),468            SelfAttention1d(mid_channels, mid_channels // 32),469            SelfAttention1d(out_channels, out_channels // 32),470        ]471 472        self.attentions = nn.ModuleList(attentions)473        self.resnets = nn.ModuleList(resnets)474 475    def forward(self, hidden_states, temb=None):476        hidden_states = self.down(hidden_states)477 478        for resnet, attn in zip(self.resnets, self.attentions):479            hidden_states = resnet(hidden_states)480            hidden_states = attn(hidden_states)481 482        return hidden_states, (hidden_states,)483 484 485class DownBlock1D(nn.Module):486    def __init__(self, out_channels, in_channels, mid_channels=None):487        super().__init__()488        mid_channels = out_channels if mid_channels is None else mid_channels489 490        self.down = Downsample1d("cubic")491        resnets = [492            ResConvBlock(in_channels, mid_channels, mid_channels),493            ResConvBlock(mid_channels, mid_channels, mid_channels),494            ResConvBlock(mid_channels, mid_channels, out_channels),495        ]496 497        self.resnets = nn.ModuleList(resnets)498 499    def forward(self, hidden_states, temb=None):500        hidden_states = self.down(hidden_states)501 502        for resnet in self.resnets:503            hidden_states = resnet(hidden_states)504 505        return hidden_states, (hidden_states,)506 507 508class DownBlock1DNoSkip(nn.Module):509    def __init__(self, out_channels, in_channels, mid_channels=None):510        super().__init__()511        mid_channels = out_channels if mid_channels is None else mid_channels512 513        resnets = [514            ResConvBlock(in_channels, mid_channels, mid_channels),515            ResConvBlock(mid_channels, mid_channels, mid_channels),516            ResConvBlock(mid_channels, mid_channels, out_channels),517        ]518 519        self.resnets = nn.ModuleList(resnets)520 521    def forward(self, hidden_states, temb=None):522        hidden_states = torch.cat([hidden_states, temb], dim=1)523        for resnet in self.resnets:524            hidden_states = resnet(hidden_states)525 526        return hidden_states, (hidden_states,)527 528 529class AttnUpBlock1D(nn.Module):530    def __init__(self, in_channels, out_channels, mid_channels=None):531        super().__init__()532        mid_channels = out_channels if mid_channels is None else mid_channels533 534        resnets = [535            ResConvBlock(2 * in_channels, mid_channels, mid_channels),536            ResConvBlock(mid_channels, mid_channels, mid_channels),537            ResConvBlock(mid_channels, mid_channels, out_channels),538        ]539        attentions = [540            SelfAttention1d(mid_channels, mid_channels // 32),541            SelfAttention1d(mid_channels, mid_channels // 32),542            SelfAttention1d(out_channels, out_channels // 32),543        ]544 545        self.attentions = nn.ModuleList(attentions)546        self.resnets = nn.ModuleList(resnets)547        self.up = Upsample1d(kernel="cubic")548 549    def forward(self, hidden_states, res_hidden_states_tuple, temb=None):550        res_hidden_states = res_hidden_states_tuple[-1]551        hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1)552 553        for resnet, attn in zip(self.resnets, self.attentions):554            hidden_states = resnet(hidden_states)555            hidden_states = attn(hidden_states)556 557        hidden_states = self.up(hidden_states)558 559        return hidden_states560 561 562class UpBlock1D(nn.Module):563    def __init__(self, in_channels, out_channels, mid_channels=None):564        super().__init__()565        mid_channels = in_channels if mid_channels is None else mid_channels566 567        resnets = [568            ResConvBlock(2 * in_channels, mid_channels, mid_channels),569            ResConvBlock(mid_channels, mid_channels, mid_channels),570            ResConvBlock(mid_channels, mid_channels, out_channels),571        ]572 573        self.resnets = nn.ModuleList(resnets)574        self.up = Upsample1d(kernel="cubic")575 576    def forward(self, hidden_states, res_hidden_states_tuple, temb=None):577        res_hidden_states = res_hidden_states_tuple[-1]578        hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1)579 580        for resnet in self.resnets:581            hidden_states = resnet(hidden_states)582 583        hidden_states = self.up(hidden_states)584 585        return hidden_states586 587 588class UpBlock1DNoSkip(nn.Module):589    def __init__(self, in_channels, out_channels, mid_channels=None):590        super().__init__()591        mid_channels = in_channels if mid_channels is None else mid_channels592 593        resnets = [594            ResConvBlock(2 * in_channels, mid_channels, mid_channels),595            ResConvBlock(mid_channels, mid_channels, mid_channels),596            ResConvBlock(mid_channels, mid_channels, out_channels, is_last=True),597        ]598 599        self.resnets = nn.ModuleList(resnets)600 601    def forward(self, hidden_states, res_hidden_states_tuple, temb=None):602        res_hidden_states = res_hidden_states_tuple[-1]603        hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1)604 605        for resnet in self.resnets:606            hidden_states = resnet(hidden_states)607 608        return hidden_states609 610 611def get_down_block(down_block_type, num_layers, in_channels, out_channels, temb_channels, add_downsample):612    if down_block_type == "DownResnetBlock1D":613        return DownResnetBlock1D(614            in_channels=in_channels,615            num_layers=num_layers,616            out_channels=out_channels,617            temb_channels=temb_channels,618            add_downsample=add_downsample,619        )620    elif down_block_type == "DownBlock1D":621        return DownBlock1D(out_channels=out_channels, in_channels=in_channels)622    elif down_block_type == "AttnDownBlock1D":623        return AttnDownBlock1D(out_channels=out_channels, in_channels=in_channels)624    elif down_block_type == "DownBlock1DNoSkip":625        return DownBlock1DNoSkip(out_channels=out_channels, in_channels=in_channels)626    raise ValueError(f"{down_block_type} does not exist.")627 628 629def get_up_block(up_block_type, num_layers, in_channels, out_channels, temb_channels, add_upsample):630    if up_block_type == "UpResnetBlock1D":631        return UpResnetBlock1D(632            in_channels=in_channels,633            num_layers=num_layers,634            out_channels=out_channels,635            temb_channels=temb_channels,636            add_upsample=add_upsample,637        )638    elif up_block_type == "UpBlock1D":639        return UpBlock1D(in_channels=in_channels, out_channels=out_channels)640    elif up_block_type == "AttnUpBlock1D":641        return AttnUpBlock1D(in_channels=in_channels, out_channels=out_channels)642    elif up_block_type == "UpBlock1DNoSkip":643        return UpBlock1DNoSkip(in_channels=in_channels, out_channels=out_channels)644    raise ValueError(f"{up_block_type} does not exist.")645 646 647def get_mid_block(mid_block_type, num_layers, in_channels, mid_channels, out_channels, embed_dim, add_downsample):648    if mid_block_type == "MidResTemporalBlock1D":649        return MidResTemporalBlock1D(650            num_layers=num_layers,651            in_channels=in_channels,652            out_channels=out_channels,653            embed_dim=embed_dim,654            add_downsample=add_downsample,655        )656    elif mid_block_type == "ValueFunctionMidBlock1D":657        return ValueFunctionMidBlock1D(in_channels=in_channels, out_channels=out_channels, embed_dim=embed_dim)658    elif mid_block_type == "UNetMidBlock1D":659        return UNetMidBlock1D(in_channels=in_channels, mid_channels=mid_channels, out_channels=out_channels)660    raise ValueError(f"{mid_block_type} does not exist.")661 662 663def get_out_block(*, out_block_type, num_groups_out, embed_dim, out_channels, act_fn, fc_dim):664    if out_block_type == "OutConv1DBlock":665        return OutConv1DBlock(num_groups_out, out_channels, embed_dim, act_fn)666    elif out_block_type == "ValueFunction":667        return OutValueFunctionBlock(fc_dim, embed_dim)668    return None669