import tensorflow as tf import keras from keras import layers, models, ops class RMSNorm(layers.Layer): def __init__(self, epsilon=1e-6, **kwargs): super().__init__(**kwargs) self.epsilon = epsilon def build(self, input_shape): self.scale = self.add_weight( name='scale', shape=(input_shape[-1],), initializer='ones', trainable=True ) def call(self, x): x_f32 = tf.cast(x, tf.float32) variance = tf.reduce_mean(tf.square(x_f32), axis=-1, keepdims=True) x_normed = x_f32 * tf.math.rsqrt(variance + self.epsilon) return x_normed * self.scale class TokenAndPositionEmbedding(layers.Layer): def __init__(self, d_model, moves=64, **kwargs): if 'position' in kwargs: kwargs.pop('position') super(TokenAndPositionEmbedding, self).__init__(**kwargs) self.d_model = d_model self.moves = moves self.row_embedding = layers.Embedding(8, d_model, name="row_emb") self.col_embedding = layers.Embedding(8, d_model, name="col_emb") self.time_embedding = layers.Embedding(moves + 1, d_model, name="time_emb") def call(self, inputs): x, board = inputs positions = tf.range(start=0, limit=64, delta=1, dtype=tf.int32) r_emb = self.row_embedding(positions // 8) c_emb = self.col_embedding(positions % 8) stone_count = tf.reduce_sum(board, axis=[1, 2, 3]) current_moves = tf.cast(stone_count, tf.int32) - 4 current_moves = tf.maximum(current_moves, 0) t_emb = self.time_embedding(current_moves) t_emb = tf.expand_dims(t_emb, axis=1) return x + tf.cast(r_emb, x.dtype) + tf.cast(c_emb, x.dtype) + tf.cast(t_emb, x.dtype) class MHA(layers.Layer): def __init__(self, d_model, num_heads, rate=0.2, use_8dir_mask=False, **kwargs): super().__init__(**kwargs) self.use_8dir_mask = use_8dir_mask self.att = layers.MultiHeadAttention(num_heads=num_heads, key_dim=d_model//num_heads) self.rmsnorm = RMSNorm() self.dropout = layers.Dropout(rate) if self.use_8dir_mask: import numpy as np mask = np.zeros((64, 64), dtype=bool) for i in range(64): r1, c1 = divmod(i, 8) for j in range(64): r2, c2 = divmod(j, 8) if r1 == r2 or c1 == c2 or abs(r1 - r2) == abs(c1 - c2): mask[i, j] = True self.attn_mask = tf.constant(mask, dtype=tf.bool) self.attn_mask = tf.reshape(self.attn_mask, (1, 1, 64, 64)) def call(self, x, training=False): x_f32 = tf.cast(x, tf.float32) normed_inputs = self.rmsnorm(x_f32) if self.use_8dir_mask: attn_output = self.att( query = normed_inputs, value = normed_inputs, key = normed_inputs, attention_mask = self.attn_mask, training = training ) else: attn_output = self.att( query = normed_inputs, value = normed_inputs, key = normed_inputs, training = training ) attn_output = self.dropout(attn_output, training=training) return x_f32 + tf.cast(attn_output, tf.float32) class FFN(layers.Layer): def __init__(self, d_model, rate=0.2, **kwargs): super().__init__(**kwargs) ff_dim = int(d_model * 8 / 3) self.w1 = layers.Dense(ff_dim, name="w1") self.w2 = layers.Dense(ff_dim, name="w2") self.w3 = layers.Dense(d_model, name="w3") self.rmsnorm = RMSNorm() self.dropout = layers.Dropout(rate) def call(self, x, training=False): x_f32 = tf.cast(x, tf.float32) normed_inputs = self.rmsnorm(x_f32) gate = tf.nn.silu(self.w1(normed_inputs)) hidden = gate * self.w2(normed_inputs) ffn_output = self.w3(hidden) ffn_output = self.dropout(ffn_output, training=training) return x_f32 + tf.cast(ffn_output, tf.float32) class DynamicAssembly(layers.Layer): def __init__(self, d_model, num_heads, num_mha=2, num_ffn=2, steps=2, rate=0.2, enable_mask=False, top_k=1, **kwargs): super().__init__(**kwargs) self.d_model = d_model self.steps = steps self.num_mha = num_mha self.num_ffn = num_ffn self.enable_mask = enable_mask self.top_k = top_k self.mha_pool = [] for i in range(num_mha): use_mask = self.enable_mask and (i % 2 == 1) self.mha_pool.append(MHA(d_model, num_heads, rate, use_8dir_mask=use_mask, name=f"pool_mha_{i}")) self.ffn_pool = [] for i in range(num_ffn): self.ffn_pool.append(FFN(d_model, rate, name=f"pool_ffn_{i}")) self.mha_router = layers.Dense(num_mha, name="mha_router") self.ffn_router = layers.Dense(num_ffn, name="ffn_router") self.step_embedding = layers.Embedding(steps, d_model) self.last_probs = [] def route_and_execute(self, x, pool, router, num_options, step_vec, training): x_pooled = tf.reduce_mean(x, axis=1) router_input = x_pooled + step_vec logits = router(router_input) probs = tf.nn.softmax(logits, axis=-1) k = min(self.top_k, num_options) if k < num_options: _, topk_indices = tf.math.top_k(probs, k=k) mask = tf.reduce_sum(tf.one_hot(topk_indices, depth=num_options), axis=1) mask = tf.cast(mask, probs.dtype) if training: dispatch_frac = tf.reduce_mean(mask, axis=0) prob_frac = tf.reduce_mean(probs, axis=0) balancing_loss = num_options * tf.reduce_sum(dispatch_frac * prob_frac) self.add_loss(tf.cast(0.01 * balancing_loss, tf.float32)) routed_probs = probs * mask routed_probs = routed_probs / (tf.reduce_sum(routed_probs, axis=-1, keepdims=True) + 1e-9) else: routed_probs = probs outputs = [layer(x, training=training) for layer in pool] stacked_outputs = tf.stack(outputs, axis=1) probs_bc = tf.expand_dims(routed_probs, axis=-1) probs_bc = tf.expand_dims(probs_bc, axis=-1) probs_bc = tf.cast(probs_bc, tf.float32) weighted_sum = tf.reduce_sum(stacked_outputs * probs_bc, axis=1) return tf.cast(weighted_sum, x.dtype), probs def call(self, x, training=False): x = tf.cast(x, tf.float32) if not training: self.last_probs = [] for i in range(self.steps): step_vec = tf.cast(self.step_embedding(tf.convert_to_tensor([i])), tf.float32) x, mha_probs = self.route_and_execute(x, self.mha_pool, self.mha_router, self.num_mha, step_vec, training) x, ffn_probs = self.route_and_execute(x, self.ffn_pool, self.ffn_router, self.num_ffn, step_vec, training) if not training: self.last_probs.append(tf.concat([mha_probs, ffn_probs], axis=-1)) return x class AttentionPooling(layers.Layer): def __init__(self, d_model, num_heads=4, **kwargs): super().__init__(**kwargs) self.d_model = d_model self.num_heads = num_heads def build(self, input_shape): self.query = self.add_weight( name='query', shape=(1, 1, self.d_model), initializer='random_normal', trainable=True ) self.mha = layers.MultiHeadAttention(num_heads=self.num_heads, key_dim=self.d_model // self.num_heads) self.rmsnorm = RMSNorm() def call(self, x, training=False): batch_size = tf.shape(x)[0] q = tf.tile(self.query, [batch_size, 1, 1]) pooled = self.mha(query=q, value=x, key=x, training=training) pooled = self.rmsnorm(pooled) return tf.squeeze(pooled, axis=1) def build_model(config): d_model = config.get('embed_dim', 128) num_blocks = config.get('block', 4) num_heads = config.get('head', 4) num_mha = config.get('num_mha', 2) num_ffn = config.get('num_ffn', 2) steps = config.get('steps', 2) dropout_rate = config.get('dropout', 0.2) enable_mask = config.get('enable_mask', False) input_shape = (8, 8, 3) inputs = layers.Input(shape=input_shape, dtype=tf.float32) x = layers.Reshape((64, 3))(inputs) x = layers.Dense(d_model)(x) x = TokenAndPositionEmbedding(d_model, 64)([x, inputs]) for _ in range(num_blocks): x = DynamicAssembly(d_model, num_heads, num_mha=num_mha, num_ffn=num_ffn, steps=steps, rate=dropout_rate, enable_mask=enable_mask)(x) # Policy Head policy_x = RMSNorm()(x) policy_x = layers.Dense(d_model, activation='relu', name="policy_hidden")(policy_x) policy_logits = layers.Dense(1, name="policy_logits")(policy_x) policy_logits = layers.Reshape((64,))(policy_logits) policy_head = layers.Activation('softmax', name='p', dtype='float32')(policy_logits) # Value Head value_x = AttentionPooling(d_model, num_heads=num_heads)(x) value_shared = layers.Dense(128, activation='relu', name="value_shared")(value_x) # V1: Win rate win_hidden = layers.Dense(64, activation='relu', name="win_hidden")(value_shared) win_out = layers.Dense(1, activation='tanh', name="win_out")(win_hidden) # V2: Score diff score_hidden = layers.Dense(64, activation='relu', name="score_hidden")(value_shared) score_out = layers.Dense(1, activation='tanh', name="score_out")(score_hidden) # V1 + V2 value_head = layers.Concatenate(name='v', axis=-1)([win_out, score_out]) return keras.Model(inputs=inputs, outputs=[policy_head, value_head], name="moe_2") if __name__ == '__main__': conf = {'embed_dim': 128, 'block': 4, 'head': 4, 'num_mha': 3, 'num_ffn': 2, 'steps': 2} # conf = {'embed_dim': 96, 'block': 3, 'head': 3, 'num_mha': 2, 'num_ffn': 2, 'steps': 1} model = build_model(conf) model.summary() print(f"Total Params: {model.count_params()}")