import tensorflow as tf import keras from keras import layers, models, ops 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, **kwargs): super().__init__(**kwargs) self.att = layers.MultiHeadAttention(num_heads=num_heads, key_dim=d_model//num_heads) self.layernorm = layers.LayerNormalization(epsilon=1e-6) self.dropout = layers.Dropout(rate) def call(self, x, training=False): x_f32 = tf.cast(x, tf.float32) normed_inputs = self.layernorm(x_f32) 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 = d_model * 4 self.ffn = models.Sequential([layers.Dense(ff_dim, activation='gelu'),layers.Dense(d_model)]) self.layernorm = layers.LayerNormalization(epsilon=1e-6) self.dropout = layers.Dropout(rate) def call(self, x, training=False): x_f32 = tf.cast(x, tf.float32) normed_inputs = self.layernorm(x_f32) ffn_output = self.ffn(normed_inputs) 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=4, num_ffn=4, steps=8, rate=0.2, **kwargs): super().__init__(**kwargs) self.d_model = d_model self.steps = steps self.num_options = num_mha + num_ffn self.layer_pool = [] for i in range(num_mha): self.layer_pool.append(MHA(d_model, num_heads, rate, name=f"pool_mha_{i}")) for i in range(num_ffn): self.layer_pool.append(FFN(d_model, rate, name=f"pool_ffn_{i}")) self.router_dense = layers.Dense(self.num_options, name="router") self.step_embedding = layers.Embedding(steps, d_model) self.last_probs = [] def call(self, x, training=False): if not training: self.last_probs = [] for i in range(self.steps): step_vec = self.step_embedding(tf.convert_to_tensor([i])) x_pooled = tf.reduce_mean(x, axis=1) router_input = x_pooled + tf.cast(step_vec, x_pooled.dtype) logits = self.router_dense(router_input) probs = tf.nn.softmax(logits, axis=-1) if not training: self.last_probs.append(probs) outputs = [layer(x, training=training) for layer in self.layer_pool] stacked_outputs = tf.stack(outputs, axis=1) probs_bc = tf.expand_dims(probs, axis=-1) probs_bc = tf.expand_dims(probs_bc, axis=-1) probs_bc = tf.cast(probs_bc, stacked_outputs.dtype) weighted_sum = tf.reduce_sum(stacked_outputs * probs_bc, axis=1) x = tf.cast(weighted_sum, x.dtype) return x def build_model(config): d_model = config.get('embed_dim', 32) num_blocks = config.get('block', 2) num_heads = config.get('head', 4) num_mha = config.get('num_mha', 2) num_ffn = config.get('num_ffn', 2) steps = config.get('steps', 3) dropout_rate = config.get('dropout', 0.2) input_shape = (8, 8, 2) inputs = layers.Input(shape=input_shape, dtype=tf.float32) x = layers.Reshape((64, 2))(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)(x) # Policy Head policy_x = layers.Dense(d_model, activation='relu', name="policy_hidden")(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 (M1.h5 actual structure: Conv1D -> Flatten -> Dense(128)) value_x = layers.Conv1D(8, 1, activation='relu')(x) value_x = layers.Flatten()(value_x) value_x = layers.Dense(128, activation='relu', name="value_hidden")(value_x) value_head = layers.Dense(1, activation='tanh', name='v', dtype='float32')(value_x) model = models.Model(inputs=inputs, outputs=[policy_head, value_head]) return model if __name__ == '__main__': conf = {'embed_dim': 128, 'block': 4, 'head': 4, 'num_mha': 2, 'num_ffn': 2, 'steps': 2} model = build_model(conf) model.summary() print(f"Total Params: {model.count_params()}")