{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":52950,"databundleVersionId":5973250,"sourceType":"competition"},{"sourceId":11694503,"sourceType":"datasetVersion","datasetId":7340011}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport shutil\nimport pyarrow.parquet as pq\nimport tensorflow as tf\nimport json\nimport matplotlib\nimport matplotlib.pyplot as plt\nimport random\nimport math # For ceil\n\nfrom skimage.transform import resize\nfrom tensorflow import keras\nfrom tensorflow.keras import layers # Sử dụng keras.layers thay vì tf.keras.layers\nfrom tqdm.notebook import tqdm\nfrom matplotlib import animation, rc\nimport numpy as np # linear algebra\nimport pandas as pd","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T12:25:14.348983Z","iopub.execute_input":"2025-05-14T12:25:14.349205Z","iopub.status.idle":"2025-05-14T12:25:28.514429Z","shell.execute_reply.started":"2025-05-14T12:25:14.349174Z","shell.execute_reply":"2025-05-14T12:25:28.513698Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define Global Constants\nFRAME_LEN = 128  # Max sequence length for landmarks\nTARGET_LEN = 64 # Max sequence length for phrases (must match model's target_maxlen)\nNUM_JOINTS = 26 # 21 hand + 5 pose\nNUM_COORDINATES = 3 # X, Y, Z\n\ndataset_df = pd.read_csv('/kaggle/input/asl-fingerspelling/train.csv')\n\n# Create a list of paths to .tfrecord files\ntf_records_all = dataset_df['file_id'].map(\n    lambda x: f'/kaggle/input/pre-data-fsp/new_data/{x}.tfrecord'\n).unique()\nprint(f\"List of {len(tf_records_all)} TFRecord files.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T12:25:28.515216Z","iopub.execute_input":"2025-05-14T12:25:28.515564Z","iopub.status.idle":"2025-05-14T12:25:28.697932Z","shell.execute_reply.started":"2025-05-14T12:25:28.515548Z","shell.execute_reply":"2025-05-14T12:25:28.697267Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"random.shuffle(tf_records_all)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T12:25:28.699326Z","iopub.execute_input":"2025-05-14T12:25:28.699539Z","iopub.status.idle":"2025-05-14T12:25:28.703142Z","shell.execute_reply.started":"2025-05-14T12:25:28.699523Z","shell.execute_reply":"2025-05-14T12:25:28.702424Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_split_ratio = 0.8\nnum_total_tf_records = len(tf_records_all)\nsplit_idx = int(num_total_tf_records * train_split_ratio)\n\ntrain_tf_records = tf_records_all[:split_idx]\nvalid_tf_records = tf_records_all[split_idx:]\n\nprint(f\"Number of training TFRecord files: {len(train_tf_records)}\")\nprint(f\"Number of validation TFRecord files: {len(valid_tf_records)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T12:25:28.703938Z","iopub.execute_input":"2025-05-14T12:25:28.704131Z","iopub.status.idle":"2025-05-14T12:25:28.720829Z","shell.execute_reply.started":"2025-05-14T12:25:28.704115Z","shell.execute_reply":"2025-05-14T12:25:28.720168Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"LPOSE = [13, 15, 17, 19, 21]\nRPOSE = [14, 16, 18, 20, 22]\nPOSE = LPOSE + RPOSE\nX_COORDS = [f'x_right_hand_{i}' for i in range(21)] + [f'x_left_hand_{i}' for i in range(21)] + [f'x_pose_{i}' for i in POSE]\nY_COORDS = [f'y_right_hand_{i}' for i in range(21)] + [f'y_left_hand_{i}' for i in range(21)] + [f'y_pose_{i}' for i in POSE]\nZ_COORDS = [f'z_right_hand_{i}' for i in range(21)] + [f'z_left_hand_{i}' for i in range(21)] + [f'z_pose_{i}' for i in POSE]\nFEATURE_COLUMNS = X_COORDS + Y_COORDS + Z_COORDS\n\nRHAND_LNDMRK_IDX = [i for i, col in enumerate(X_COORDS) if \"right_hand\" in col]\nLHAND_LNDMRK_IDX = [i for i, col in enumerate(X_COORDS) if \"left_hand\" in col]\nRPOSE_LNDMRK_IDX = [i for i, col in enumerate(X_COORDS) if \"pose\" in col and int(col.split('_')[-1]) in RPOSE]\nLPOSE_LNDMRK_IDX = [i for i, col in enumerate(X_COORDS) if \"pose\" in col and int(col.split('_')[-1]) in LPOSE]\n\n# Corrected landmark indices\nRHAND_OFFSET_Y = len(X_COORDS)\nRHAND_OFFSET_Z = len(X_COORDS) + len(Y_COORDS)\nLHAND_OFFSET_Y = len(X_COORDS)\nLHAND_OFFSET_Z = len(X_COORDS) + len(Y_COORDS)\nRPOSE_OFFSET_Y = len(X_COORDS)\nRPOSE_OFFSET_Z = len(X_COORDS) + len(Y_COORDS)\nLPOSE_OFFSET_Y = len(X_COORDS)\nLPOSE_OFFSET_Z = len(X_COORDS) + len(Y_COORDS)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T12:25:28.721497Z","iopub.execute_input":"2025-05-14T12:25:28.721703Z","iopub.status.idle":"2025-05-14T12:25:28.736261Z","shell.execute_reply.started":"2025-05-14T12:25:28.721688Z","shell.execute_reply":"2025-05-14T12:25:28.735592Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"RHAND_IDX = np.concatenate([\n    RHAND_LNDMRK_IDX,\n    np.array(RHAND_LNDMRK_IDX) + RHAND_OFFSET_Y,\n    np.array(RHAND_LNDMRK_IDX) + RHAND_OFFSET_Z\n]).astype(int)\nLHAND_IDX = np.concatenate([\n    LHAND_LNDMRK_IDX,\n    np.array(LHAND_LNDMRK_IDX) + LHAND_OFFSET_Y,\n    np.array(LHAND_LNDMRK_IDX) + LHAND_OFFSET_Z\n]).astype(int)\nRPOSE_IDX = np.concatenate([\n    RPOSE_LNDMRK_IDX,\n    np.array(RPOSE_LNDMRK_IDX) + RPOSE_OFFSET_Y,\n    np.array(RPOSE_LNDMRK_IDX) + RPOSE_OFFSET_Z\n]).astype(int)\nLPOSE_IDX = np.concatenate([\n    LPOSE_LNDMRK_IDX,\n    np.array(LPOSE_LNDMRK_IDX) + LPOSE_OFFSET_Y,\n    np.array(LPOSE_LNDMRK_IDX) + LPOSE_OFFSET_Z\n]).astype(int)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T12:25:28.736821Z","iopub.execute_input":"2025-05-14T12:25:28.737036Z","iopub.status.idle":"2025-05-14T12:25:28.758469Z","shell.execute_reply.started":"2025-05-14T12:25:28.737012Z","shell.execute_reply":"2025-05-14T12:25:28.757815Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"with open (\"/kaggle/input/asl-fingerspelling/character_to_prediction_index.json\", \"r\") as f:\n    char_to_num = json.load(f)\n\npad_token = 'P'\nstart_token = '<'\nend_token = '>'\npad_token_idx = 59\nstart_token_idx = 60\nend_token_idx = 61\n\nchar_to_num[pad_token] = pad_token_idx\nchar_to_num[start_token] = start_token_idx\nchar_to_num[end_token] = end_token_idx\nnum_to_char = {j:i for i,j in char_to_num.items()}\nVOCAB_SIZE = len(char_to_num)\n\nprint(f\"Vocabulary size: {VOCAB_SIZE}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T12:25:28.759096Z","iopub.execute_input":"2025-05-14T12:25:28.759323Z","iopub.status.idle":"2025-05-14T12:25:28.778867Z","shell.execute_reply.started":"2025-05-14T12:25:28.759307Z","shell.execute_reply":"2025-05-14T12:25:28.778307Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def resize_pad(x, frame_len, num_joints, num_coords):\n    current_frames = tf.shape(x)[0]\n    if current_frames < frame_len:\n        padding_frames = frame_len - current_frames\n        x = tf.pad(x, ([[0, padding_frames], [0, 0], [0, 0]]))\n    else:\n        x = tf.image.resize(x, (frame_len, x.shape[1])) # Use x.shape[1] for num_joints\n    return x\n\ndef pre_process(x):\n    rhand_all_coords = tf.gather(x, RHAND_IDX, axis=1)\n    lhand_all_coords = tf.gather(x, LHAND_IDX, axis=1)\n    rpose_all_coords = tf.gather(x, RPOSE_IDX, axis=1)\n    lpose_all_coords = tf.gather(x, LPOSE_IDX, axis=1)\n\n    current_frames_count = tf.shape(x)[0]\n    rhand_proc = tf.reshape(rhand_all_coords, (current_frames_count, 21, 3))\n    lhand_proc = tf.reshape(lhand_all_coords, (current_frames_count, 21, 3))\n    rpose_proc = tf.reshape(rpose_all_coords, (current_frames_count, 5, 3))\n    lpose_proc = tf.reshape(lpose_all_coords, (current_frames_count, 5, 3))\n\n    rnan_frames = tf.reduce_any(tf.math.is_nan(rhand_proc), axis=[1,2])\n    lnan_frames = tf.reduce_any(tf.math.is_nan(lhand_proc), axis=[1,2])\n\n    rnans_count = tf.math.count_nonzero(rnan_frames)\n    lnans_count = tf.math.count_nonzero(lnan_frames)\n\n    if rnans_count > lnans_count:\n        hand = lhand_proc\n        pose = lpose_proc\n        hand = tf.concat([1.0 - hand[..., 0:1], hand[..., 1:3]], axis=-1)\n        pose = tf.concat([1.0 - pose[..., 0:1], pose[..., 1:3]], axis=-1)\n    else:\n        hand = rhand_proc\n        pose = rpose_proc\n\n    mean = tf.math.reduce_mean(hand, axis=1, keepdims=True)\n    std = tf.math.reduce_std(hand, axis=1, keepdims=True)\n    hand = (hand - mean) / (std + 1e-6)\n\n    processed_landmarks = tf.concat([hand, pose], axis=1)\n    processed_landmarks = resize_pad(processed_landmarks, FRAME_LEN, NUM_JOINTS, NUM_COORDINATES)\n    processed_landmarks = tf.where(tf.math.is_nan(processed_landmarks), tf.zeros_like(processed_landmarks), processed_landmarks)\n    processed_landmarks = tf.reshape(processed_landmarks, (FRAME_LEN, NUM_JOINTS * NUM_COORDINATES))\n    return processed_landmarks\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T12:25:28.779571Z","iopub.execute_input":"2025-05-14T12:25:28.779775Z","iopub.status.idle":"2025-05-14T12:25:28.794564Z","shell.execute_reply.started":"2025-05-14T12:25:28.779759Z","shell.execute_reply":"2025-05-14T12:25:28.793996Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def decode_fn(record_bytes):\n    schema = {COL: tf.io.VarLenFeature(dtype=tf.float32) for COL in FEATURE_COLUMNS}\n    schema[\"phrase\"] = tf.io.FixedLenFeature([], dtype=tf.string)\n    features = tf.io.parse_single_example(record_bytes, schema)\n    phrase = features[\"phrase\"]\n    landmarks_sparse = [features[COL] for COL in FEATURE_COLUMNS]\n    landmarks_dense = [tf.sparse.to_dense(s) for s in landmarks_sparse]\n\n    for i in range(len(landmarks_dense)):\n        if tf.shape(landmarks_dense[i])[0] == 0:\n            landmarks_dense[i] = tf.zeros((1,1), dtype=tf.float32)\n        landmarks_dense[i] = tf.reshape(landmarks_dense[i], [-1])\n\n    landmarks = tf.transpose(landmarks_dense)\n    if tf.shape(landmarks)[0] == 0:\n        landmarks = tf.zeros((1, len(FEATURE_COLUMNS)), dtype=tf.float32)\n    return landmarks, phrase\n\ntable = tf.lookup.StaticHashTable(\n    initializer=tf.lookup.KeyValueTensorInitializer(\n        keys=list(char_to_num.keys()),\n        values=list(char_to_num.values()),\n    ),\n    default_value=tf.constant(char_to_num.get(' ', pad_token_idx)),\n    name=\"char_to_int_lookup\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T12:25:28.796873Z","iopub.execute_input":"2025-05-14T12:25:28.797103Z","iopub.status.idle":"2025-05-14T12:25:30.114613Z","shell.execute_reply.started":"2025-05-14T12:25:28.797088Z","shell.execute_reply":"2025-05-14T12:25:30.113875Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def convert_fn(landmarks, phrase):\n    phrase_bytes = tf.strings.join([tf.constant(start_token.encode('utf-8')), phrase, tf.constant(end_token.encode('utf-8'))])\n    phrase_split = tf.strings.bytes_split(phrase_bytes)\n    phrase_tokens = table.lookup(phrase_split)\n    phrase_padded = tf.pad(phrase_tokens, paddings=[[0, TARGET_LEN - tf.shape(phrase_tokens)[0]]],\n                           mode='CONSTANT', constant_values=pad_token_idx)\n    processed_landmarks = pre_process(landmarks)\n    return processed_landmarks, phrase_padded\n\n\nBATCH_SIZE = 32\n\ntrain_ds = tf.data.TFRecordDataset(train_tf_records, num_parallel_reads=tf.data.AUTOTUNE) \\\n    .map(decode_fn, num_parallel_calls=tf.data.AUTOTUNE) \\\n    .map(convert_fn, num_parallel_calls=tf.data.AUTOTUNE) \\\n    .batch(BATCH_SIZE) \\\n    .prefetch(buffer_size=tf.data.AUTOTUNE) \\\n    .cache() \\\n    .repeat()\n\nvalid_ds = tf.data.TFRecordDataset(valid_tf_records, num_parallel_reads=tf.data.AUTOTUNE) \\\n    .map(decode_fn, num_parallel_calls=tf.data.AUTOTUNE) \\\n    .map(convert_fn, num_parallel_calls=tf.data.AUTOTUNE) \\\n    .batch(BATCH_SIZE) \\\n    .prefetch(buffer_size=tf.data.AUTOTUNE) \\\n    .cache()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T12:25:30.115337Z","iopub.execute_input":"2025-05-14T12:25:30.115598Z","iopub.status.idle":"2025-05-14T12:25:36.446056Z","shell.execute_reply.started":"2025-05-14T12:25:30.11558Z","shell.execute_reply":"2025-05-14T12:25:36.445477Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# MODEL","metadata":{}},{"cell_type":"code","source":"# --- Model Definition ---\nclass TokenEmbedding(layers.Layer):\n    def __init__(self, num_vocab, maxlen, num_hid): # Bỏ giá trị mặc định\n        super().__init__()\n        self.emb = layers.Embedding(num_vocab, num_hid)\n        self.pos_emb = layers.Embedding(input_dim=maxlen, output_dim=num_hid)\n\n    def call(self, x):\n        maxlen = tf.shape(x)[-1]\n        x = self.emb(x)\n        positions = tf.range(start=0, limit=maxlen, delta=1)\n        positions = self.pos_emb(positions)\n        return x + positions\n\nclass STTRBlock(layers.Layer):\n    def __init__(self, num_hid, num_heads, dropout_rate): # Bỏ giá trị mặc định\n        super().__init__()\n        key_dim = num_hid // num_heads\n        self.spatial_att = layers.MultiHeadAttention(num_heads=num_heads, key_dim=key_dim, dropout=dropout_rate)\n        self.temporal_att = layers.MultiHeadAttention(num_heads=num_heads, key_dim=key_dim, dropout=dropout_rate)\n        self.norm1_spatial = layers.LayerNormalization(epsilon=1e-6)\n        self.norm2_spatial = layers.LayerNormalization(epsilon=1e-6)\n        self.ffn_spatial = keras.Sequential([\n            layers.Dense(num_hid * 2, activation=\"relu\"),\n            layers.Dense(num_hid),\n            layers.Dropout(dropout_rate)\n        ])\n        self.norm1_temporal = layers.LayerNormalization(epsilon=1e-6)\n        self.norm2_temporal = layers.LayerNormalization(epsilon=1e-6)\n        self.ffn_temporal = keras.Sequential([\n            layers.Dense(num_hid * 2, activation=\"relu\"),\n            layers.Dense(num_hid),\n            layers.Dropout(dropout_rate)\n        ])\n\n    def call(self, x, training=False):\n        B, T, J, D_hid = tf.shape(x)[0], tf.shape(x)[1], tf.shape(x)[2], tf.shape(x)[3]\n        h_spatial = tf.reshape(x, [B * T, J, D_hid])\n        att_spatial_input = self.norm1_spatial(h_spatial)\n        att_spatial_out = self.spatial_att(query=att_spatial_input, value=att_spatial_input, key=att_spatial_input, training=training)\n        h_spatial = h_spatial + att_spatial_out\n        ffn_spatial_input = self.norm2_spatial(h_spatial)\n        ffn_spatial_out = self.ffn_spatial(ffn_spatial_input, training=training)\n        h_spatial = h_spatial + ffn_spatial_out\n        h_spatial_reshaped = tf.reshape(h_spatial, [B, T, J, D_hid])\n        h_temporal = tf.transpose(h_spatial_reshaped, [0, 2, 1, 3])\n        h_temporal = tf.reshape(h_temporal, [B * J, T, D_hid])\n        att_temporal_input = self.norm1_temporal(h_temporal)\n        att_temporal_out = self.temporal_att(query=att_temporal_input, value=att_temporal_input, key=att_temporal_input, training=training)\n        h_temporal = h_temporal + att_temporal_out\n        ffn_temporal_input = self.norm2_temporal(h_temporal)\n        ffn_temporal_out = self.ffn_temporal(ffn_temporal_input, training=training)\n        h_temporal = h_temporal + ffn_temporal_out\n        h_temporal_reshaped = tf.reshape(h_temporal, [B, J, T, D_hid])\n        out = tf.transpose(h_temporal_reshaped, [0, 2, 1, 3])\n        return out","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T12:25:36.446734Z","iopub.execute_input":"2025-05-14T12:25:36.447001Z","iopub.status.idle":"2025-05-14T12:25:36.457001Z","shell.execute_reply.started":"2025-05-14T12:25:36.446978Z","shell.execute_reply":"2025-05-14T12:25:36.456329Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class STTRModel(keras.Model): # Kế thừa từ keras.Model\n    def __init__(self, num_joints, input_coord_dim, num_hid, num_blocks, num_heads, dropout_rate): # Bỏ giá trị mặc định\n        super().__init__()\n        self.input_proj = layers.Dense(num_hid, activation='relu')\n        self.time_pos_emb = layers.Embedding(input_dim=FRAME_LEN, output_dim=num_hid)\n        self.sttr_blocks = [STTRBlock(num_hid=num_hid, num_heads=num_heads, dropout_rate=dropout_rate) for _ in range(num_blocks)]\n        self.dropout = layers.Dropout(dropout_rate)\n\n    def call(self, x, training=False):\n        B, T, J, C_in = tf.shape(x)[0], tf.shape(x)[1], tf.shape(x)[2], tf.shape(x)[3]\n        x = self.input_proj(x)\n        positions_time = tf.range(start=0, limit=T, delta=1)\n        time_embed = self.time_pos_emb(positions_time)\n        x = x + time_embed[tf.newaxis, :, tf.newaxis, :]\n        x = self.dropout(x, training=training)\n        for block in self.sttr_blocks:\n            x = block(x, training=training)\n        x = tf.reduce_mean(x, axis=2)\n        return x\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T12:25:36.457774Z","iopub.execute_input":"2025-05-14T12:25:36.458522Z","iopub.status.idle":"2025-05-14T12:25:36.477446Z","shell.execute_reply.started":"2025-05-14T12:25:36.458498Z","shell.execute_reply":"2025-05-14T12:25:36.476866Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class TransformerDecoderLayer(layers.Layer):\n    def __init__(self, embed_dim, num_heads, feed_forward_dim, dropout_rate): # Bỏ giá trị mặc định\n        super().__init__()\n        self.self_att = layers.MultiHeadAttention(\n            num_heads=num_heads, key_dim=embed_dim, dropout=dropout_rate\n        )\n        self.enc_att = layers.MultiHeadAttention(\n            num_heads=num_heads, key_dim=embed_dim, dropout=dropout_rate\n        )\n        self.ffn = keras.Sequential([\n            layers.Dense(feed_forward_dim, activation=\"relu\"),\n            layers.Dense(embed_dim),\n            layers.Dropout(dropout_rate) # Thêm dropout ở FFN\n        ])\n        self.norm1 = layers.LayerNormalization(epsilon=1e-6)\n        self.norm2 = layers.LayerNormalization(epsilon=1e-6)\n        self.norm3 = layers.LayerNormalization(epsilon=1e-6)\n        self.dropout1 = layers.Dropout(dropout_rate)\n        self.dropout2 = layers.Dropout(dropout_rate)\n        self.dropout3 = layers.Dropout(dropout_rate) # Giữ lại FFN dropout\n\n    def call(self, target, enc_out, training, look_ahead_mask=None, padding_mask=None):\n        normed_target = self.norm1(target)\n        attn1 = self.self_att(query=normed_target, value=normed_target, key=normed_target,\n                              attention_mask=look_ahead_mask, training=training)\n        attn1 = self.dropout1(attn1, training=training)\n        out1 = target + attn1\n        normed_out1 = self.norm2(out1)\n        attn2 = self.enc_att(query=normed_out1, value=enc_out, key=enc_out,\n                             attention_mask=padding_mask, training=training)\n        attn2 = self.dropout2(attn2, training=training)\n        out2 = out1 + attn2\n        normed_out2 = self.norm3(out2)\n        ffn_output = self.ffn(normed_out2, training=training)\n        # ffn_output = self.dropout3(ffn_output, training=training) # Dropout đã có trong self.ffn\n        out3 = out2 + ffn_output\n        return out3\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T12:25:36.478172Z","iopub.execute_input":"2025-05-14T12:25:36.478392Z","iopub.status.idle":"2025-05-14T12:25:36.496663Z","shell.execute_reply.started":"2025-05-14T12:25:36.478371Z","shell.execute_reply":"2025-05-14T12:25:36.496028Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class TransformerDecoder(keras.layers.Layer):\n    def __init__(self, num_layers, embed_dim, num_heads, feed_forward_dim, dropout_rate): # Bỏ giá trị mặc định\n        super().__init__()\n        self.num_layers = num_layers\n        self.dec_layers = [\n            TransformerDecoderLayer(embed_dim, num_heads, feed_forward_dim, dropout_rate)\n            for _ in range(num_layers)\n        ]\n    def call(self, target, enc_out, training, look_ahead_mask, padding_mask):\n        x = target\n        for i in range(self.num_layers):\n            x = self.dec_layers[i](\n                target=x,\n                enc_out=enc_out,\n                training=training,\n                look_ahead_mask=look_ahead_mask,\n                padding_mask=padding_mask\n            )\n        return x\n\n\nclass Transformer(keras.Model):\n    def __init__(\n        self,\n        num_hid, # Bỏ giá trị mặc định\n        num_head,\n        num_feed_forward,\n        source_maxlen,\n        target_maxlen,\n        num_layers_enc_sttr,\n        num_layers_dec,\n        num_classes,\n        dropout_rate,\n    ):\n        super().__init__()\n        # Các metrics tùy chỉnh sẽ được tạo trong self.compile() hoặc build()\n        # self.loss_tracker = keras.metrics.Mean(name=\"loss\")\n        # self.accuracy_tracker = keras.metrics.Mean(name=\"accuracy\") # Sẽ dùng masked_accuracy\n        # Các metrics không được cập nhật này nên được loại bỏ hoặc xử lý đúng cách\n        # self.acc_metric = keras.metrics.Mean(name=\"edit_dist_mean_placeholder\")\n        # self.edit_distance_tracker = keras.metrics.MeanMetricWrapper(fn=tf.edit_distance, name=\"edit_distance_raw_placeholder\")\n\n\n        self.target_maxlen = target_maxlen\n        self.num_classes = num_classes\n\n        self.encoder = STTRModel(\n            num_joints=NUM_JOINTS,\n            input_coord_dim=NUM_COORDINATES,\n            num_hid=num_hid,\n            num_blocks=num_layers_enc_sttr,\n            num_heads=num_head,\n            dropout_rate=dropout_rate\n        )\n        self.token_embedding = TokenEmbedding(\n            num_vocab=num_classes, maxlen=target_maxlen, num_hid=num_hid\n        )\n        self.decoder = TransformerDecoder(\n            num_layers=num_layers_dec,\n            embed_dim=num_hid,\n            num_heads=num_head,\n            feed_forward_dim=num_feed_forward,\n            dropout_rate=dropout_rate\n        )\n        self.final_classifier = layers.Dense(num_classes)\n        self.dropout_output = layers.Dropout(dropout_rate)\n\n    def _create_look_ahead_mask(self, size):\n        mask = 1 - tf.linalg.band_part(tf.ones((size, size)), -1, 0)\n        return mask\n\n    def call(self, inputs, training=False):\n        source_flat, target_tokens = inputs\n        B = tf.shape(source_flat)[0]\n        source = tf.reshape(source_flat, [B, FRAME_LEN, NUM_JOINTS, NUM_COORDINATES])\n        if training:\n            noise = tf.random.normal(shape=tf.shape(source), mean=0.0, stddev=0.01)\n            source = source + noise\n        enc_out = self.encoder(source, training=training)\n        target_embedded = self.token_embedding(target_tokens)\n        look_ahead_mask = self._create_look_ahead_mask(tf.shape(target_tokens)[1])\n        dec_out = self.decoder(target_embedded, enc_out, training=training,\n                               look_ahead_mask=look_ahead_mask,\n                               padding_mask=None)\n        dec_out = self.dropout_output(dec_out, training=training)\n        final_output = self.final_classifier(dec_out)\n        return final_output\n\n    # Keras 3: compute_loss, train_step, test_step, metrics được quản lý tốt hơn\n    # Ghi đè train_step và test_step theo Keras 3 conventions\n\n    def train_step(self, data):\n        # Unpack the data. Its structure depends on your model and\n        # on what you pass to `fit()`.\n        # Trong trường hợp này, dataset trả về (source, target)\n        source_flat, target_full = data\n\n        dec_input = target_full[:, :-1]\n        dec_target = target_full[:, 1:] # y_true cho loss và metrics\n        loss_mask = tf.math.logical_not(tf.math.equal(dec_target, pad_token_idx))\n        effective_sample_weight = tf.cast(loss_mask, tf.float32)\n\n\n        with tf.GradientTape() as tape:\n            y_pred = self([source_flat, dec_input], training=True)  # Forward pass, y_pred là logits\n            # Compute the loss value.\n            # The `compile()` method A) configures the loss calculator B) configures the metrics calculator.\n            # A) `self.loss` is the loss calculator.\n            loss = self.loss(y_true=dec_target, y_pred=y_pred, sample_weight=effective_sample_weight)\n            # Keras 3: self.loss đã bao gồm regularization losses nếu loss được compile là object\n            # Nếu loss_fn là function, bạn cần thêm self.losses thủ công\n            if isinstance(self.compiled_loss, keras.losses.Loss): # Check if compiled_loss is an object\n                 pass # Regularization losses are handled by the Loss object's __call__\n            elif self.losses: # Nếu loss là một hàm, thêm regularization losses\n                 loss += tf.add_n(self.losses)\n\n\n        # Compute gradients\n        trainable_vars = self.trainable_variables\n        gradients = tape.gradient(loss, trainable_vars)\n        # Update weights\n        self.optimizer.apply_gradients(zip(gradients, trainable_vars))\n        # Update the metrics.\n        # Metrics are configured in `compile()`.\n        # B) `self.metrics` is the metrics calculator.\n        # Keras 3:\n        for metric in self.metrics: # self.metrics ở đây là các metrics đã compile\n            if metric.name == \"loss\": # Đây là compiled loss metric, không phải loss scalar ở trên\n                metric.update_state(loss) # Cập nhật compiled loss metric với scalar loss\n            else: # Các metrics khác (ví dụ: masked_accuracy)\n                metric.update_state(dec_target, y_pred, sample_weight=effective_sample_weight)\n        # Return a dict mapping metric names to current value.\n        return {m.name: m.result() for m in self.metrics}\n\n\n    def test_step(self, data):\n        source_flat, target_full = data\n        dec_input = target_full[:, :-1]\n        dec_target = target_full[:, 1:]\n        loss_mask = tf.math.logical_not(tf.math.equal(dec_target, pad_token_idx))\n        effective_sample_weight = tf.cast(loss_mask, tf.float32)\n\n        y_pred = self([source_flat, dec_input], training=False)\n        # Updates the metrics using data from the validation dataset.\n        loss = self.loss(y_true=dec_target, y_pred=y_pred, sample_weight=effective_sample_weight)\n        if isinstance(self.compiled_loss, keras.losses.Loss):\n             pass\n        elif self.losses:\n             loss += tf.add_n(self.losses)\n\n        for metric in self.metrics:\n            if metric.name == \"loss\":\n                metric.update_state(loss)\n            else:\n                metric.update_state(dec_target, y_pred, sample_weight=effective_sample_weight)\n        return {m.name: m.result() for m in self.metrics}\n\n\n    def generate(self, source_flat, target_start_token_idx):\n        bs = tf.shape(source_flat)[0]\n        source_reshaped = tf.reshape(source_flat, [bs, FRAME_LEN, NUM_JOINTS, NUM_COORDINATES])\n        enc_out = self.encoder(source_reshaped, training=False)\n        dec_input_tokens = tf.ones((bs, 1), dtype=tf.int32) * target_start_token_idx\n        for _ in range(self.target_maxlen -1):\n            target_embedded = self.token_embedding(dec_input_tokens)\n            look_ahead_mask = self._create_look_ahead_mask(tf.shape(dec_input_tokens)[1])\n            dec_out = self.decoder(target_embedded, enc_out, training=False,\n                                   look_ahead_mask=look_ahead_mask, padding_mask=None)\n            last_token_logits = self.final_classifier(dec_out[:, -1:, :])\n            next_token = tf.argmax(last_token_logits, axis=-1, output_type=tf.int32)\n            dec_input_tokens = tf.concat([dec_input_tokens, next_token], axis=1)\n            if tf.reduce_all(tf.reduce_any(tf.equal(dec_input_tokens, end_token_idx), axis=1)):\n                 break\n        return dec_input_tokens","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T12:25:36.497446Z","iopub.execute_input":"2025-05-14T12:25:36.498141Z","iopub.status.idle":"2025-05-14T12:25:36.520014Z","shell.execute_reply.started":"2025-05-14T12:25:36.498124Z","shell.execute_reply":"2025-05-14T12:25:36.519304Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DisplayOutputs(keras.callbacks.Callback):\n    def __init__(\n        self, batch_data_iterator, idx_to_token_map, target_start_token_idx, target_end_token_idx, num_samples_to_display=3\n    ):\n        super().__init__()\n        self.batch_data_iterator = batch_data_iterator\n        self.idx_to_char = idx_to_token_map\n        self.target_start_token_idx = target_start_token_idx\n        self.target_end_token_idx = target_end_token_idx\n        self.num_samples_to_display = num_samples_to_display\n\n    def on_epoch_end(self, epoch, logs=None):\n        if (epoch + 1) % 4 != 0:\n            return\n        try:\n            source_batch, target_batch = next(self.batch_data_iterator)\n        except StopIteration: # Reset iterator nếu hết batch\n            self.batch_data_iterator = iter(valid_ds_for_callback) # Cần định nghĩa valid_ds_for_callback ở scope ngoài\n            source_batch, target_batch = next(self.batch_data_iterator)\n\n        preds_tokens = self.model.generate(source_batch, self.target_start_token_idx)\n\n        print(f\"\\n--- Epoch {epoch+1} Sample Predictions ---\")\n        for i in range(min(tf.shape(source_batch)[0], self.num_samples_to_display)):\n            target_text_full = \"\".join([self.idx_to_char.get(t_id, '?') for t_id in target_batch[i].numpy()])\n            cleaned_target_text = target_text_full.replace(self.idx_to_char.get(self.target_start_token_idx, ''), '') \\\n                                                 .replace(self.idx_to_char.get(self.target_end_token_idx, ''), '') \\\n                                                 .replace(self.idx_to_char.get(pad_token_idx, ''), '')\n            prediction_text = \"\"\n            for token_id in preds_tokens[i].numpy():\n                char = self.idx_to_char.get(token_id, '?')\n                if token_id == self.target_end_token_idx:\n                    break\n                if token_id != self.target_start_token_idx and token_id != pad_token_idx :\n                    prediction_text += char\n            print(f\"Target    : {cleaned_target_text}\")\n            print(f\"Prediction: {prediction_text}\\n\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T12:25:36.520825Z","iopub.execute_input":"2025-05-14T12:25:36.521044Z","iopub.status.idle":"2025-05-14T12:25:36.542527Z","shell.execute_reply.started":"2025-05-14T12:25:36.521023Z","shell.execute_reply":"2025-05-14T12:25:36.541824Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"valid_ds_for_callback = tf.data.TFRecordDataset(valid_tf_records, num_parallel_reads=tf.data.AUTOTUNE) \\\n    .map(decode_fn, num_parallel_calls=tf.data.AUTOTUNE) \\\n    .map(convert_fn, num_parallel_calls=tf.data.AUTOTUNE) \\\n    .batch(BATCH_SIZE) \\\n    .prefetch(buffer_size=tf.data.AUTOTUNE)\ndisplay_cb_iterator = iter(valid_ds_for_callback)\n\n\ndisplay_cb = DisplayOutputs(\n    display_cb_iterator,\n    idx_to_token_map=num_to_char,\n    target_start_token_idx=start_token_idx,\n    target_end_token_idx=end_token_idx,\n    num_samples_to_display=min(BATCH_SIZE, 5) # Hiển thị ít hơn để không quá dài\n)\n\nearlystop_cb = keras.callbacks.EarlyStopping(\n    monitor='val_loss',\n    patience=10, # Tăng patience\n    restore_best_weights=True,\n    verbose=1\n)\n\nmodel = Transformer(\n    num_hid=200,\n    num_head=4,\n    num_feed_forward=256,\n    source_maxlen=FRAME_LEN,\n    target_maxlen=TARGET_LEN,\n    num_layers_enc_sttr=2,\n    num_layers_dec=2,\n    num_classes=VOCAB_SIZE,\n    dropout_rate=0.2\n)\n\n# Sử dụng reduction mặc định của Keras cho loss_fn\n# Keras sẽ tự động xử lý sample_weight để bỏ qua padding\nloss_fn_object = keras.losses.SparseCategoricalCrossentropy(from_logits=True)\n# Nếu bạn muốn reduction=\"none\" và tính mean thủ công, bạn phải cẩn thận trong train_step\n\noptimizer = keras.optimizers.Adam(learning_rate=1e-4)\n\ndef masked_accuracy(y_true, y_pred_logits): # y_pred_logits là output của model\n    # y_true đã là dec_target (tức là target_full[:, 1:])\n    # y_pred_logits có shape (batch_size, seq_len_pred, num_classes)\n    # seq_len_pred thường là TARGET_LEN - 1\n    \n    mask = tf.math.logical_not(tf.math.equal(y_true, pad_token_idx))\n    y_pred_tokens = tf.argmax(y_pred_logits, axis=-1, output_type=y_true.dtype) # Cùng dtype với y_true\n    \n    matches = tf.cast(tf.equal(y_true, y_pred_tokens), dtype=tf.float32)\n    masked_matches = matches * tf.cast(mask, tf.float32)\n    \n    total_masked_elements = tf.reduce_sum(tf.cast(mask, tf.float32))\n    # Tránh chia cho 0 nếu tất cả đều là padding (mặc dù hiếm)\n    return tf.math.divide_no_nan(tf.reduce_sum(masked_matches), total_masked_elements)\n\n\nmodel.compile(optimizer=optimizer, loss=loss_fn_object, metrics=[masked_accuracy])\n\nsteps_per_epoch = 632 # Nên tính toán dựa trên kích thước dataset\nvalidation_steps = 164 # Nên tính toán\n\nprint(f\"Using steps_per_epoch: {steps_per_epoch}, validation_steps: {validation_steps}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T12:25:36.543192Z","iopub.execute_input":"2025-05-14T12:25:36.543423Z","iopub.status.idle":"2025-05-14T12:25:40.745323Z","shell.execute_reply.started":"2025-05-14T12:25:36.543404Z","shell.execute_reply":"2025-05-14T12:25:40.744659Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"history = model.fit(\n    train_ds,\n    validation_data=valid_ds,\n    epochs=50,\n    steps_per_epoch=steps_per_epoch,\n    validation_steps=validation_steps,\n    callbacks=[display_cb, earlystop_cb]\n)\n\nprint(\"Training finished.\")\nmodel.save_weights(\"asl_transformer_sttr_weights.keras\") # Sử dụng định dạng .keras mới\nprint(\"Model weights saved.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T12:25:40.746037Z","iopub.execute_input":"2025-05-14T12:25:40.74631Z"}},"outputs":[],"execution_count":null}]}