{"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}],"dockerImageVersionId":31012,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install mediapipe","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T00:25:10.617304Z","iopub.execute_input":"2025-05-19T00:25:10.617559Z","iopub.status.idle":"2025-05-19T00:25:27.254921Z","shell.execute_reply.started":"2025-05-19T00:25:10.617534Z","shell.execute_reply":"2025-05-19T00:25:27.254015Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport shutil\nimport numpy as np\nimport pandas as pd\nimport pyarrow.parquet as pq\nimport tensorflow as tf\nimport json\nimport mediapipe\nimport matplotlib\nimport matplotlib.pyplot as plt\nimport random\n\nfrom skimage.transform import resize\nfrom mediapipe.framework.formats import landmark_pb2\nfrom tensorflow import keras\nfrom tensorflow.keras import layers\nfrom tqdm.notebook import tqdm\nfrom matplotlib import animation, rc\nfrom tensorflow.keras.callbacks import EarlyStopping\n\ndataset_df = pd.read_csv('/kaggle/input/asl-fingerspelling/train.csv')\nprint(\"Full train dataset shape is {}\".format(dataset_df.shape))\n# Fetch sequence_id, file_id, phrase from first row\nsequence_id, file_id, phrase = dataset_df.iloc[0][['sequence_id', 'file_id', 'phrase']]\nprint(f\"sequence_id: {sequence_id}, file_id: {file_id}, phrase: {phrase}\")\n# Fetch data from parquet file\nsample_sequence_df = pq.read_table(f\"/kaggle/input/asl-fingerspelling/train_landmarks/{str(file_id)}.parquet\",\n    filters=[[('sequence_id', '=', sequence_id)],]).to_pandas()\nprint(\"Full sequence dataset shape is {}\".format(sample_sequence_df.shape))\nmp_pose = mediapipe.solutions.pose\nmp_hands = mediapipe.solutions.hands\nmp_drawing = mediapipe.solutions.drawing_utils \nmp_drawing_styles = mediapipe.solutions.drawing_styles\n\ndef get_hands(seq_df):\n    images = []\n    all_hand_landmarks = []\n    for seq_idx in range(len(seq_df)):\n        x_hand = seq_df.iloc[seq_idx].filter(regex=\"x_right_hand.*\").values\n        y_hand = seq_df.iloc[seq_idx].filter(regex=\"y_right_hand.*\").values\n        z_hand = seq_df.iloc[seq_idx].filter(regex=\"z_right_hand.*\").values\n\n        right_hand_image = np.zeros((600, 600, 3))\n\n        right_hand_landmarks = landmark_pb2.NormalizedLandmarkList()\n        \n        for x, y, z in zip(x_hand, y_hand, z_hand):\n            right_hand_landmarks.landmark.add(x=x, y=y, z=z)\n\n        mp_drawing.draw_landmarks(\n                right_hand_image,\n                right_hand_landmarks,\n                mp_hands.HAND_CONNECTIONS,\n                landmark_drawing_spec=mp_drawing_styles.get_default_hand_landmarks_style())\n        \n        x_hand = seq_df.iloc[seq_idx].filter(regex=\"x_left_hand.*\").values\n        y_hand = seq_df.iloc[seq_idx].filter(regex=\"y_left_hand.*\").values\n        z_hand = seq_df.iloc[seq_idx].filter(regex=\"z_left_hand.*\").values\n        \n        left_hand_image = np.zeros((600, 600, 3))\n        \n        left_hand_landmarks = landmark_pb2.NormalizedLandmarkList()\n        for x, y, z in zip(x_hand, y_hand, z_hand):\n            left_hand_landmarks.landmark.add(x=x, y=y, z=z)\n\n        mp_drawing.draw_landmarks(\n                left_hand_image,\n                left_hand_landmarks,\n                mp_hands.HAND_CONNECTIONS,\n                landmark_drawing_spec=mp_drawing_styles.get_default_hand_landmarks_style())\n        \n        images.append([right_hand_image.astype(np.uint8), left_hand_image.astype(np.uint8)])\n        all_hand_landmarks.append([right_hand_landmarks, left_hand_landmarks])\n    return images, all_hand_landmarks\n# Get the images created using mediapipe apis\nhand_images, hand_landmarks = get_hands(sample_sequence_df)\n# Pose coordinates for hand movement.\nLPOSE = [13, 15, 17, 19, 21]\nRPOSE = [14, 16, 18, 20, 22]\nPOSE = LPOSE + RPOSE\nX = [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 = [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 = [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 + Y + Z\nX_IDX = [i for i, col in enumerate(FEATURE_COLUMNS)  if \"x_\" in col]\nY_IDX = [i for i, col in enumerate(FEATURE_COLUMNS)  if \"y_\" in col]\nZ_IDX = [i for i, col in enumerate(FEATURE_COLUMNS)  if \"z_\" in col]\n\nRHAND_IDX = [i for i, col in enumerate(FEATURE_COLUMNS)  if \"right\" in col]\nLHAND_IDX = [i for i, col in enumerate(FEATURE_COLUMNS)  if  \"left\" in col]\nRPOSE_IDX = [i for i, col in enumerate(FEATURE_COLUMNS)  if  \"pose\" in col and int(col[-2:]) in RPOSE]\nLPOSE_IDX = [i for i, col in enumerate(FEATURE_COLUMNS)  if  \"pose\" in col and int(col[-2:]) in LPOSE]\n# Set length of frames to 128\nFRAME_LEN = 128\n\n# Create directory to store the new data\nif not os.path.isdir(\"preprocessed\"):\n    os.mkdir(\"preprocessed\")\nelse:\n    shutil.rmtree(\"preprocessed\")\n    os.mkdir(\"preprocessed\")\n\n# Loop through each file_id\nfor file_id in tqdm(dataset_df.file_id.unique()):\n    # Parquet file name\n    pq_file = f\"/kaggle/input/asl-fingerspelling/train_landmarks/{file_id}.parquet\"\n    # Filter train.csv and fetch entries only for the relevant file_id\n    file_df = dataset_df.loc[dataset_df[\"file_id\"] == file_id]\n    # Fetch the parquet file\n    parquet_df = pq.read_table(f\"/kaggle/input/asl-fingerspelling/train_landmarks/{str(file_id)}.parquet\",\n                              columns=['sequence_id'] + FEATURE_COLUMNS).to_pandas()\n    # File name for the updated data\n    tf_file = f\"preprocessed/{file_id}.tfrecord\"\n    parquet_numpy = parquet_df.to_numpy()\n    # Initialize the pointer to write the output of \n    # each `for loop` below as a sequence into the file.\n    with tf.io.TFRecordWriter(tf_file) as file_writer:\n        # Loop through each sequence in file.\n        for seq_id, phrase in zip(file_df.sequence_id, file_df.phrase):\n            # Fetch sequence data\n            frames = parquet_numpy[parquet_df.index == seq_id]\n            \n            # Calculate the number of NaN values in each hand landmark\n            r_nonan = np.sum(np.sum(np.isnan(frames[:, RHAND_IDX]), axis = 1) == 0)\n            l_nonan = np.sum(np.sum(np.isnan(frames[:, LHAND_IDX]), axis = 1) == 0)\n            no_nan = max(r_nonan, l_nonan)\n            \n            if 2*len(phrase)<no_nan:\n                features = {FEATURE_COLUMNS[i]: tf.train.Feature(\n                    float_list=tf.train.FloatList(value=frames[:, i])) for i in range(len(FEATURE_COLUMNS))}\n                features[\"phrase\"] = tf.train.Feature(bytes_list=tf.train.BytesList(value=[bytes(phrase, 'utf-8')]))\n                record_bytes = tf.train.Example(features=tf.train.Features(feature=features)).SerializeToString()\n                file_writer.write(record_bytes)\ntf_records = dataset_df.file_id.map(lambda x: f'/kaggle/working/preprocessed/{x}.tfrecord').unique()\nprint(f\"List of {len(tf_records)} TFRecord files.\")\nwith open (\"/kaggle/input/asl-fingerspelling/character_to_prediction_index.json\", \"r\") as f:\n    char_to_num = json.load(f)\n\n# Add pad_token, start pointer and end pointer to the dict\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()}\n# Function to resize and add padding.\ndef resize_pad(x):\n    if tf.shape(x)[0] < FRAME_LEN:\n        x = tf.pad(x, ([[0, FRAME_LEN-tf.shape(x)[0]], [0, 0], [0, 0]]))\n    else:\n        x = tf.image.resize(x, (FRAME_LEN, tf.shape(x)[1]))\n    return x\n\n# Detect the dominant hand from the number of NaN values.\n# Dominant hand will have less NaN values since it is in frame moving.\ndef pre_process(x):\n    rhand = tf.gather(x, RHAND_IDX, axis=1)\n    lhand = tf.gather(x, LHAND_IDX, axis=1)\n    rpose = tf.gather(x, RPOSE_IDX, axis=1)\n    lpose = tf.gather(x, LPOSE_IDX, axis=1)\n    \n    rnan_idx = tf.reduce_any(tf.math.is_nan(rhand), axis=1)\n    lnan_idx = tf.reduce_any(tf.math.is_nan(lhand), axis=1)\n    \n    rnans = tf.math.count_nonzero(rnan_idx)\n    lnans = tf.math.count_nonzero(lnan_idx)\n    \n    # For dominant hand\n    if rnans > lnans:\n        hand = lhand\n        pose = lpose\n        \n        hand_x = hand[:, 0*(len(LHAND_IDX)//3) : 1*(len(LHAND_IDX)//3)]\n        hand_y = hand[:, 1*(len(LHAND_IDX)//3) : 2*(len(LHAND_IDX)//3)]\n        hand_z = hand[:, 2*(len(LHAND_IDX)//3) : 3*(len(LHAND_IDX)//3)]\n        hand = tf.concat([1-hand_x, hand_y, hand_z], axis=1)\n        \n        pose_x = pose[:, 0*(len(LPOSE_IDX)//3) : 1*(len(LPOSE_IDX)//3)]\n        pose_y = pose[:, 1*(len(LPOSE_IDX)//3) : 2*(len(LPOSE_IDX)//3)]\n        pose_z = pose[:, 2*(len(LPOSE_IDX)//3) : 3*(len(LPOSE_IDX)//3)]\n        pose = tf.concat([1-pose_x, pose_y, pose_z], axis=1)\n    else:\n        hand = rhand\n        pose = rpose\n    \n    hand_x = hand[:, 0*(len(LHAND_IDX)//3) : 1*(len(LHAND_IDX)//3)]\n    hand_y = hand[:, 1*(len(LHAND_IDX)//3) : 2*(len(LHAND_IDX)//3)]\n    hand_z = hand[:, 2*(len(LHAND_IDX)//3) : 3*(len(LHAND_IDX)//3)]\n    hand = tf.concat([hand_x[..., tf.newaxis], hand_y[..., tf.newaxis], hand_z[..., tf.newaxis]], axis=-1)\n    \n    mean = tf.math.reduce_mean(hand, axis=1)[:, tf.newaxis, :]\n    std = tf.math.reduce_std(hand, axis=1)[:, tf.newaxis, :]\n    hand = (hand - mean) / std\n\n    pose_x = pose[:, 0*(len(LPOSE_IDX)//3) : 1*(len(LPOSE_IDX)//3)]\n    pose_y = pose[:, 1*(len(LPOSE_IDX)//3) : 2*(len(LPOSE_IDX)//3)]\n    pose_z = pose[:, 2*(len(LPOSE_IDX)//3) : 3*(len(LPOSE_IDX)//3)]\n    pose = tf.concat([pose_x[..., tf.newaxis], pose_y[..., tf.newaxis], pose_z[..., tf.newaxis]], axis=-1)\n    \n    x = tf.concat([hand, pose], axis=1)\n    x = resize_pad(x)\n    \n    x = tf.where(tf.math.is_nan(x), tf.zeros_like(x), x)\n    x = tf.reshape(x, (FRAME_LEN, len(LHAND_IDX) + len(LPOSE_IDX)))\n    return x\n\ndef 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 = ([tf.sparse.to_dense(features[COL]) for COL in FEATURE_COLUMNS])\n    # Transpose to maintain the original shape of landmarks data.\n    landmarks = tf.transpose(landmarks)\n    \n    return landmarks, phrase\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(-1),\n    name=\"class_weight\"\n)\n\ndef convert_fn(landmarks, phrase):\n    # Add start and end pointers to phrase.\n    phrase = start_token + phrase + end_token\n    phrase = tf.strings.bytes_split(phrase)\n    phrase = table.lookup(phrase)\n    # Vectorize and add padding.\n    phrase = tf.pad(phrase, paddings=[[0, 64 - tf.shape(phrase)[0]]], mode = 'CONSTANT',\n                    constant_values = pad_token_idx)\n    # Apply pre_process function to the landmarks.\n    return pre_process(landmarks), phrase\n\nbatch_size = 64\ntrain_len = int(0.8 * len(tf_records))\n\ntrain_ds = tf.data.TFRecordDataset(tf_records[:train_len]).map(decode_fn).map(convert_fn).batch(batch_size).prefetch(buffer_size=tf.data.AUTOTUNE).cache()\nvalid_ds = tf.data.TFRecordDataset(tf_records[train_len:]).map(decode_fn).map(convert_fn).batch(batch_size).prefetch(buffer_size=tf.data.AUTOTUNE).cache()\nclass TokenEmbedding(layers.Layer):\n    def __init__(self, num_vocab=1000, maxlen=100, num_hid=64):\n        super().__init__()\n        self.emb = tf.keras.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\n\nclass LandmarkEmbedding(layers.Layer):\n    def __init__(self, num_hid=64, maxlen=100):\n        super().__init__()\n        self.conv1 = tf.keras.layers.Conv1D(\n            num_hid, 11, strides=2, padding=\"same\", activation=\"relu\"\n        )\n        self.conv2 = tf.keras.layers.Conv1D(\n            num_hid, 11, strides=2, padding=\"same\", activation=\"relu\"\n        )\n        self.conv3 = tf.keras.layers.Conv1D(\n            num_hid, 11, strides=2, padding=\"same\", activation=\"relu\"\n        )\n        self.pos_emb = layers.Embedding(input_dim=maxlen, output_dim=num_hid)\n\n    def call(self, x):\n        x = self.conv1(x)\n        x = self.conv2(x)\n        return self.conv3(x)\nclass TransformerEncoder(layers.Layer):\n    def __init__(self, embed_dim, num_heads, feed_forward_dim, rate=0.1):\n        super().__init__()\n        self.att = layers.MultiHeadAttention(num_heads=num_heads, key_dim=embed_dim)\n        self.ffn = keras.Sequential(\n            [\n                layers.Dense(feed_forward_dim, activation=\"relu\"),\n                layers.Dense(embed_dim),\n            ]\n        )\n        self.layernorm1 = layers.LayerNormalization(epsilon=1e-6)\n        self.layernorm2 = layers.LayerNormalization(epsilon=1e-6)\n        self.dropout1 = layers.Dropout(rate)\n        self.dropout2 = layers.Dropout(rate)\n\n    def call(self, inputs, training=None):\n        attn_output = self.att(inputs, inputs)\n        attn_output = self.dropout1(attn_output, training=training)\n        out1 = self.layernorm1(inputs + attn_output)\n        ffn_output = self.ffn(out1)\n        ffn_output = self.dropout2(ffn_output, training=training)\n        return self.layernorm2(out1 + ffn_output)\nclass TransformerDecoder(layers.Layer):\n    def __init__(self, embed_dim, num_heads, feed_forward_dim, dropout_rate=0.1):\n        super().__init__()\n        self.layernorm1 = layers.LayerNormalization(epsilon=1e-6)\n        self.layernorm2 = layers.LayerNormalization(epsilon=1e-6)\n        self.layernorm3 = layers.LayerNormalization(epsilon=1e-6)\n        self.self_att = layers.MultiHeadAttention(\n            num_heads=num_heads, key_dim=embed_dim\n        )\n        self.enc_att = layers.MultiHeadAttention(num_heads=num_heads, key_dim=embed_dim)\n        self.self_dropout = layers.Dropout(0.5)\n        self.enc_dropout = layers.Dropout(0.1)\n        self.ffn_dropout = layers.Dropout(0.1)\n        self.ffn = keras.Sequential(\n            [\n                layers.Dense(feed_forward_dim, activation=\"relu\"),\n                layers.Dense(embed_dim),\n            ]\n        )\n\n    def causal_attention_mask(self, batch_size, n_dest, n_src, dtype):\n        \"\"\"Masks the upper half of the dot product matrix in self attention.\n\n        This prevents flow of information from future tokens to current token.\n        1's in the lower triangle, counting from the lower right corner.\n        \"\"\"\n        i = tf.range(n_dest)[:, None]\n        j = tf.range(n_src)\n        m = i >= j - n_src + n_dest\n        mask = tf.cast(m, dtype)\n        mask = tf.reshape(mask, [1, n_dest, n_src])\n        mult = tf.concat(\n            [batch_size[..., tf.newaxis], tf.constant([1, 1], dtype=tf.int32)], 0\n        )\n        return tf.tile(mask, mult)\n\n    def call(self, enc_out, target, training):\n        input_shape = tf.shape(target)\n        batch_size = input_shape[0]\n        seq_len = input_shape[1]\n        causal_mask = self.causal_attention_mask(batch_size, seq_len, seq_len, tf.bool)\n        target_att = self.self_att(target, target, attention_mask=causal_mask)\n        target_norm = self.layernorm1(target + self.self_dropout(target_att, training = training))\n        enc_out = self.enc_att(target_norm, enc_out)\n        enc_out_norm = self.layernorm2(self.enc_dropout(enc_out, training = training) + target_norm)\n        ffn_out = self.ffn(enc_out_norm)\n        ffn_out_norm = self.layernorm3(enc_out_norm + self.ffn_dropout(ffn_out, training = training))\n        return ffn_out_norm\nclass Transformer(keras.Model):\n    def __init__(\n        self,\n        num_hid=64,\n        num_head=2,\n        num_feed_forward=128,\n        source_maxlen=100,\n        target_maxlen=100,\n        num_layers_enc=4,\n        num_layers_dec=1,\n        num_classes=60,\n    ):\n        super().__init__()\n        self.loss_metric = keras.metrics.Mean(name=\"loss\")\n        self.acc_metric = keras.metrics.Mean(name=\"edit_dist\")\n        self.num_layers_enc = num_layers_enc\n        self.num_layers_dec = num_layers_dec\n        self.target_maxlen = target_maxlen\n        self.num_classes = num_classes\n\n        self.enc_input = LandmarkEmbedding(num_hid=num_hid, maxlen=source_maxlen)\n        self.dec_input = TokenEmbedding(\n            num_vocab=num_classes, maxlen=target_maxlen, num_hid=num_hid\n        )\n\n        self.encoder = keras.Sequential(\n            [self.enc_input]\n            + [\n                TransformerEncoder(num_hid, num_head, num_feed_forward)\n                for _ in range(num_layers_enc)\n            ]\n        )\n\n        for i in range(num_layers_dec):\n            setattr(\n                self,\n                f\"dec_layer_{i}\",\n                TransformerDecoder(num_hid, num_head, num_feed_forward),\n            )\n\n        self.classifier = layers.Dense(num_classes)\n\n    def decode(self, enc_out, target, training=None):\n        y = self.dec_input(target)\n        for i in range(self.num_layers_dec):\n            y = getattr(self, f\"dec_layer_{i}\")(enc_out, y, training=training)\n        return y\n\n    def call(self, inputs, training=None):\n        source = inputs[0]\n        target = inputs[1]\n        x = self.encoder(source, training=training)\n        y = self.decode(x, target, training=training)\n        return self.classifier(y)\n    def beam_search_decode(self, source, start_token, end_token, beam_width=3):\n        batch_size = tf.shape(source)[0]\n        results = []\n    \n        for i in range(batch_size):\n            src = tf.expand_dims(source[i], axis=0)\n            enc = self.encoder(src, training=False)\n    \n            sequences = [[start_token]]\n            scores = [0.0]\n    \n            for _ in range(self.target_maxlen):\n                all_candidates = []\n    \n                for i_seq in range(len(sequences)):\n                    seq = sequences[i_seq]\n                    score = scores[i_seq]\n    \n                    # Nếu đã có token kết thúc, không mở rộng nữa\n                    if seq[-1] == end_token:\n                        all_candidates.append((score, seq))\n                        continue\n    \n                    dec_input = tf.constant(seq, dtype=tf.int32)[tf.newaxis, :]\n                    dec_out = self.decode(enc, dec_input, training=False)\n                    logits = self.classifier(dec_out)\n                    log_probs = tf.nn.log_softmax(logits[:, -1, :])\n    \n                    top_k_log_probs, top_k_ids = tf.math.top_k(log_probs, k=beam_width)\n    \n                    for j in range(beam_width):\n                        candidate = seq + [int(top_k_ids[0][j])]\n                        candidate_score = score + float(top_k_log_probs[0][j])\n                        all_candidates.append((candidate_score, candidate))\n    \n                # Sắp xếp các ứng viên theo điểm số (từ cao đến thấp)\n                ordered = sorted(all_candidates, key=lambda tup: tup[0], reverse=True)\n                sequences = [cand[1] for cand in ordered[:beam_width]]\n                scores = [cand[0] for cand in ordered[:beam_width]]\n    \n                # Dừng nếu tất cả chuỗi đều đã chứa token kết thúc\n                if all(seq[-1] == end_token for seq in sequences):\n                    break\n    \n            # Lấy chuỗi tốt nhất từ beam\n            results.append(sequences[0])\n    \n        return tf.constant(results, dtype=tf.int32)\n\n    @property\n    def metrics(self):\n        return [self.loss_metric]\n\n    def train_step(self, batch):\n        \"\"\"Processes one batch inside model.fit().\"\"\"\n        source = batch[0]\n        target = batch[1]\n\n        input_shape = tf.shape(target)\n        batch_size = input_shape[0]\n        \n        dec_input = target[:, :-1]\n        dec_target = target[:, 1:]\n        with tf.GradientTape() as tape:\n            preds = self([source, dec_input], training=True)\n            one_hot = tf.one_hot(dec_target, depth=self.num_classes)\n            mask = tf.math.logical_not(tf.math.equal(dec_target, pad_token_idx))\n            loss = self.compiled_loss(one_hot, preds, sample_weight=mask)\n        trainable_vars = self.trainable_variables\n        gradients = tape.gradient(loss, trainable_vars)\n        self.optimizer.apply_gradients(zip(gradients, trainable_vars))\n        # Computes the Levenshtein distance between sequences since the evaluation\n        # metric for this contest is the normalized total levenshtein distance.\n        edit_dist = tf.edit_distance(\n            tf.cast(tf.sparse.from_dense(target), tf.int64),\n            tf.cast(tf.sparse.from_dense(tf.argmax(preds, axis=-1)), tf.int64),\n            normalize=True\n        )\n\n        edit_dist = tf.reduce_mean(edit_dist)\n        self.acc_metric.update_state(edit_dist)\n        self.loss_metric.update_state(loss)\n        return {\"loss\": self.loss_metric.result(), \"edit_dist\": self.acc_metric.result()}\n\n    def test_step(self, batch):        \n        source = batch[0]\n        target = batch[1]\n\n        input_shape = tf.shape(target)\n        batch_size = input_shape[0]\n        \n        dec_input = target[:, :-1]\n        dec_target = target[:, 1:]\n        preds = self([source, dec_input], training=False)\n        one_hot = tf.one_hot(dec_target, depth=self.num_classes)\n        mask = tf.math.logical_not(tf.math.equal(dec_target, pad_token_idx))\n        loss = self.compiled_loss(one_hot, preds, sample_weight=mask)\n        # Computes the Levenshtein distance between sequences since the evaluation\n        # metric for this contest is the normalized total levenshtein distance.\n        edit_dist = tf.edit_distance(tf.sparse.from_dense(target), \n                                     tf.sparse.from_dense(tf.cast(tf.argmax(preds, axis=1), tf.int32)))\n        edit_dist = tf.reduce_mean(edit_dist)\n        self.acc_metric.update_state(edit_dist)\n        self.loss_metric.update_state(loss)\n        return {\"loss\": self.loss_metric.result(), \"edit_dist\": self.acc_metric.result()}\n\n    def generate(self, source, target_start_token_idx, target_end_token_idx=None, use_beam_search=False, beam_width=3):\n        bs = tf.shape(source)[0]\n        if use_beam_search:\n            assert bs == 1, \"Beam search chỉ hỗ trợ batch_size = 1.\"\n            return tf.constant(\n                self.beam_search_decode(source, target_start_token_idx, target_end_token_idx, beam_width),\n                dtype=tf.int32\n            )\n    \n        # Greedy decoding\n        enc = self.encoder(source, training=False)\n        dec_input = tf.ones((bs, 1), dtype=tf.int32) * target_start_token_idx\n        for _ in range(self.target_maxlen - 1):\n            dec_out = self.decode(enc, dec_input, training=False)\n            logits = self.classifier(dec_out)\n            logits = tf.argmax(logits, axis=-1, output_type=tf.int32)\n            last_logit = logits[:, -1][..., tf.newaxis]\n            dec_input = tf.concat([dec_input, last_logit], axis=-1)\n            \n            # Kiểm tra xem tất cả các chuỗi trong batch đã chứa token kết thúc chưa\n            if target_end_token_idx is not None and tf.reduce_all(tf.equal(last_logit, target_end_token_idx)):\n                break\n    \n        return dec_input\n\n\nclass DisplayOutputs(keras.callbacks.Callback):\n    def __init__(\n        self, batch, idx_to_token, target_start_token_idx=60, target_end_token_idx=61\n    ):\n        self.batch = batch\n        self.target_start_token_idx = target_start_token_idx\n        self.target_end_token_idx = target_end_token_idx\n        self.idx_to_char = idx_to_token\n\n    def on_epoch_end(self, epoch, logs=None):\n        if epoch % 4 != 0:\n            return\n        source, target = self.batch\n        preds = self.model.generate(source, self.target_start_token_idx)\n\n        source = source.numpy()\n        target = target.numpy()\n        preds = preds.numpy()\n        bs = source.shape[0]\n\n        for i in range(bs):\n            tgt_str = \"\".join([self.idx_to_char[idx] for idx in target[i] if idx < len(self.idx_to_char)])\n            pred_str = \"\"\n            for idx in preds[i]:\n                if idx >= len(self.idx_to_char):\n                    continue\n                if idx == self.target_end_token_idx:\n                    break\n                pred_str += self.idx_to_char[idx]\n            print(f\"target:     {tgt_str.replace('-', '')}\")\n            print(f\"prediction: {pred_str}\\n\")\n\nbatch = next(iter(valid_ds))\n\n# The vocabulary to convert predicted indices into characters\nidx_to_char = list(char_to_num.keys())\ndisplay_cb = DisplayOutputs(\n    batch, idx_to_char, target_start_token_idx=char_to_num['<'], target_end_token_idx=char_to_num['>']\n)  # set the arguments as per vocabulary index for '<' and '>'\n\nmodel = Transformer(\n    num_hid=200,\n    num_head=4,\n    num_feed_forward=400,\n    source_maxlen = FRAME_LEN,\n    target_maxlen=64,\n    num_layers_enc=2,\n    num_layers_dec=1,\n    num_classes=62\n)\nloss_fn = tf.keras.losses.CategoricalCrossentropy(\n    from_logits=True, label_smoothing=0.1,\n)\n\n\noptimizer = keras.optimizers.Adam(0.0001)\nmodel.compile(\n    optimizer=tf.keras.optimizers.Adam(learning_rate=1e-4),\n    loss=loss_fn,\n    run_eagerly=True  # hoặc True để debug\n)\n\nsmall_train_ds = train_ds.take(5)  # Lấy 5 batch từ train_ds\nsmall_valid_ds = valid_ds.take(2)  # Lấy 2 batch từ valid_ds\n\n# Định nghĩa EPOCHS\n# EPOCHS = 3 # Chạy thử ít epoch thôi\n\n# print(\"Starting a small test run...\")\n# history = model.fit(\n#     small_train_ds,\n#     validation_data=small_valid_ds,\n#     callbacks=[display_cb],\n#     epochs=EPOCHS, # Sử dụng số epoch nhỏ để thử\n# )\n# print(\"Small test run finished.\")\n\n# Nếu chạy thành công, bạn có thể chạy với toàn bộ dữ liệu:\nearly_stopping_cb = EarlyStopping(\n    monitor='val_loss',  # Theo dõi validation loss\n    patience=3,          # Số epochs chờ đợi không có cải thiện trước khi dừng\n    restore_best_weights=True, # Khôi phục trọng số từ epoch có val_loss tốt nhất\n    verbose=1            # In thông báo khi dừng\n)\n\nprint(\"Starting full training run...\")\nEPOCHS_FULL = 100 # Ví dụ, đặt một số lớn\nhistory_full = model.fit(\n    train_ds,\n    validation_data=valid_ds,\n    callbacks=[display_cb, early_stopping_cb], # Thêm early_stopping_cb vào đây\n    epochs=EPOCHS_FULL,\n)\nprint(\"Full training run finished.\")\n# --- THÊM CODE VẼ ĐỒ THỊ VÀO ĐÂY ---\nif history_full and history_full.history:\n    print(\"\\nPlotting training history...\")\n    # Lấy các giá trị từ history\n    train_loss = history_full.history['loss']\n    val_loss = history_full.history.get('val_loss', None) # Dùng .get để tránh lỗi nếu không có val_loss\n    train_edit_dist = history_full.history['edit_dist']\n    val_edit_dist = history_full.history.get('val_edit_dist', None)\n    epochs_range = range(1, len(train_loss) + 1)\n\n    plt.figure(figsize=(15, 6))\n\n    # Đồ thị Loss\n    plt.subplot(1, 2, 1)\n    plt.plot(epochs_range, train_loss, label='Training Loss')\n    if val_loss:\n        plt.plot(epochs_range, val_loss, label='Validation Loss')\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.title('Training and Validation Loss')\n    plt.legend(loc='upper right')\n    plt.grid(True)\n\n    # Đồ thị Edit Distance\n    plt.subplot(1, 2, 2)\n    plt.plot(epochs_range, train_edit_dist, label='Training Edit Distance')\n    if val_edit_dist:\n        plt.plot(epochs_range, val_edit_dist, label='Validation Edit Distance')\n    plt.xlabel('Epoch')\n    plt.ylabel('Edit Distance')\n    plt.title('Training and Validation Edit Distance')\n    plt.legend(loc='lower right') # Có thể 'upper right' tùy vào dữ liệu\n    plt.grid(True)\n\n    plt.tight_layout() # Điều chỉnh layout cho đẹp\n    try:\n        plt.savefig(\"training_history.png\") # Lưu đồ thị thành file ảnh\n        print(\"Training history plot saved to training_history.png\")\n    except Exception as e:\n        print(f\"Could not save training history plot: {e}\")\n    plt.show() # Hiển thị đồ thị (nếu môi trường hỗ trợ)\nelse:\n    print(\"No training history found to plot.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T09:12:13.72548Z","iopub.execute_input":"2025-05-19T09:12:13.726127Z","execution_failed":"2025-05-19T09:12:48.92Z"}},"outputs":[],"execution_count":null}]}