haiku实现TemplatePairStack类

TemplatePairStack是实现蛋白质结构模版pair_act特征表示的类:
通过layer_stack.layer_stack(c.num_block)(block) 堆叠c.num_block(配置文件中为2)block 函数,每个block对输入pair_act 和 pair_mask执行计算流程:TriangleAttention ---> dropout ->TriangleAttention ---> dropout -> TriangleMultiplication ---> dropout -> TriangleMultiplication ---> dropout -> Transition

复制代码
import haiku as hk


class TemplatePairStack(hk.Module):
  """Pair stack for the templates.

  Jumper et al. (2021) Suppl. Alg. 16 "TemplatePairStack"
  """

  def __init__(self, config, global_config, name='template_pair_stack'):
    super().__init__(name=name)
    self.config = config
    self.global_config = global_config

  def __call__(self, pair_act, pair_mask, is_training, safe_key=None):
    """Builds TemplatePairStack module.

    Arguments:
      pair_act: Pair activations for single template, shape [N_res, N_res, c_t].
      pair_mask: Pair mask, shape [N_res, N_res].
      is_training: Whether the module is in training mode.
      safe_key: Safe key object encapsulating the random number generation key.

    Returns:
      Updated pair_act, shape [N_res, N_res, c_t].
    """

    if safe_key is None:
      safe_key = prng.SafeKey(hk.next_rng_key())

    gc = self.global_config
    c = self.config

    if not c.num_block:
      return pair_act

    def block(x):
      """One block of the template pair stack."""
      pair_act, safe_key = x

      dropout_wrapper_fn = functools.partial(
          dropout_wrapper, is_training=is_training, global_config=gc)

      safe_key, *sub_keys = safe_key.split(6)
      sub_keys = iter(sub_keys)

      pair_act = dropout_wrapper_fn(
          TriangleAttention(c.triangle_attention_starting_node, gc,
                            name='triangle_attention_starting_node'),
          pair_act,
          pair_mask,
          next(sub_keys))
      pair_act = dropout_wrapper_fn(
          TriangleAttention(c.triangle_attention_ending_node, gc,
                            name='triangle_attention_ending_node'),
          pair_act,
          pair_mask,
          next(sub_keys))
      pair_act = dropout_wrapper_fn(
          TriangleMultiplication(c.triangle_multiplication_outgoing, gc,
                                 name='triangle_multiplication_outgoing'),
          pair_act,
          pair_mask,
          next(sub_keys))
      pair_act = dropout_wrapper_fn(
          TriangleMultiplication(c.triangle_multiplication_incoming, gc,
                                 name='triangle_multiplication_incoming'),
          pair_act,
          pair_mask,
          next(sub_keys))
      pair_act = dropout_wrapper_fn(
          Transition(c.pair_transition, gc, name='pair_transition'),
          pair_act,
          pair_mask,
          next(sub_keys))

      return pair_act, safe_key

    if gc.use_remat:
      block = hk.remat(block)

    res_stack = layer_stack.layer_stack(c.num_block)(block)
    pair_act, safe_key = res_stack((pair_act, safe_key))
    return pair_act
相关推荐
moxiaoran57533 分钟前
Codex接入MCP
人工智能
天远Date Lab6 分钟前
微服务架构实战:基于天远学历信息高级版构建自动化人才准入网关
人工智能·微服务·架构·自动化
凉茶钱11 分钟前
【吃透C++】万字解析类和对象
开发语言·c++
small_wh1te_coder12 分钟前
字节技术总监30讲 AI课:3概率、信息论与损失函数|从Logits到Cross-Entropy带你吃透大模型训练
人工智能
云表无代码开发13 分钟前
日本开发者彻底破防:中文太强了!换成英文直接裂开
大数据·服务器·人工智能·microsoft·信息可视化
致Great25 分钟前
不止自动写论文!谷歌 ScientistTwo 让 AI 自己做实验、补消融、回审稿
人工智能·深度学习·机器学习
XMAIPC_Robot26 分钟前
CODESYS 实时控制 + RK182X 大模型算力扩展|RK3576 工业边缘控制器设计
人工智能·fpga开发·机器人·rk3588+fpga
迪飞特科技28 分钟前
【无标题】
android·人工智能·本地化大模型
白杨尚青33 分钟前
C++入门篇(十):string(上)——认识string:构造与三大遍历(一条龙讲透operator[]、迭代器、auto、范围for)
java·开发语言·c++·笔记·stl
anda010936 分钟前
A2UI 协议: AI 直接画界面,而不是只会打字
人工智能·ai编程