{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.8.16","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"tpu1vmV38","dataSources":[{"sourceId":46105,"databundleVersionId":5087314,"sourceType":"competition"},{"sourceId":5158870,"sourceType":"datasetVersion","datasetId":2997548},{"sourceId":5194802,"sourceType":"datasetVersion","datasetId":3020507},{"sourceId":5594542,"sourceType":"datasetVersion","datasetId":3218684}],"dockerImageVersionId":30475,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install -q tensorflow==2.12.0\n!pip install -q tensorflow-addons==0.20.0\n!pip install -q git+https://github.com/hoyso48/tf-utils@main\n!pip install --upgrade scipy\n!pip install tf_utils","metadata":{"id":"IpEQKDrDqAFP","outputId":"2affcd98-7204-4d46-c030-ce8daefe9ad4","trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:46:53.328293Z","iopub.execute_input":"2025-04-22T15:46:53.328528Z","iopub.status.idle":"2025-04-22T15:47:15.348051Z","shell.execute_reply.started":"2025-04-22T15:46:53.328507Z","shell.execute_reply":"2025-04-22T15:47:15.346841Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport tensorflow as tf\nimport tensorflow_addons as tfa\nimport matplotlib.pyplot as plt\nimport matplotlib as mpl\nimport tensorflow.keras.mixed_precision as mixed_precision\n\nfrom tqdm.autonotebook import tqdm\nimport sklearn\n\nfrom tf_utils.schedules import OneCycleLR, ListedLR\nfrom tf_utils.callbacks import Snapshot, SWA\nfrom tf_utils.learners import FGM, AWP\n\nimport os\nimport time\nimport pickle\nimport math\nimport random\nimport sys\nimport cv2\nimport gc\nimport glob\nimport datetime","metadata":{"id":"wD7tqFC_qAFQ","outputId":"426ff58f-01d7-451e-ea99-7b00405d8781","trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:47:20.993327Z","iopub.execute_input":"2025-04-22T15:47:20.993760Z","iopub.status.idle":"2025-04-22T15:48:02.014355Z","shell.execute_reply.started":"2025-04-22T15:47:20.993728Z","shell.execute_reply":"2025-04-22T15:48:02.013540Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Seed all random number generators\ndef seed_everything(seed=42):\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    random.seed(seed)\n    np.random.seed(seed)\n    tf.random.set_seed(seed)\n    \ndef get_strategy(device='TPU-VM'):\n    if \"TPU\" in device:\n        tpu = 'local' if device=='TPU-VM' else None\n        print(\"connecting to TPU...\")\n        tpu = tf.distribute.cluster_resolver.TPUClusterResolver.connect(tpu=tpu)\n        strategy = tf.distribute.TPUStrategy(tpu)\n        IS_TPU = True\n\n    if device == \"GPU\"  or device==\"CPU\":\n        ngpu = len(tf.config.experimental.list_physical_devices('GPU'))\n        if ngpu>1:\n            print(\"Using multi GPU\")\n            strategy = tf.distribute.MirroredStrategy()\n        elif ngpu==1:\n            print(\"Using single GPU\")\n            strategy = tf.distribute.get_strategy()\n        else:\n            print(\"Using CPU\")\n            strategy = tf.distribute.get_strategy()\n            CFG.device = \"CPU\"\n\n    if device == \"GPU\":\n        print(\"Num GPUs Available: \", ngpu)\n\n    AUTO     = tf.data.experimental.AUTOTUNE\n    REPLICAS = strategy.num_replicas_in_sync\n    print(f'REPLICAS: {REPLICAS}')\n    \n    return strategy, REPLICAS, IS_TPU\n\nSTRATEGY, N_REPLICAS, IS_TPU = get_strategy()","metadata":{"id":"u74o98JxqAFQ","outputId":"af042ef9-76e7-40a2-9938-e5fd4bb6f495","trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:48:30.330948Z","iopub.execute_input":"2025-04-22T15:48:30.331948Z","iopub.status.idle":"2025-04-22T15:48:38.425313Z","shell.execute_reply.started":"2025-04-22T15:48:30.331910Z","shell.execute_reply":"2025-04-22T15:48:38.424378Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import re\ndef count_data_items(filenames):\n    n = [int(re.compile(r\"-([0-9]*)\\.\").search(filename.split('/')[-1]).group(1)) for filename in filenames]\n    return np.sum(n)","metadata":{"id":"iAc9FKztqAFQ","outputId":"ade0413c-5a17-4a4f-f5a6-b0048ab76ffa","trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:49:18.606115Z","iopub.execute_input":"2025-04-22T15:49:18.606927Z","iopub.status.idle":"2025-04-22T15:49:18.611391Z","shell.execute_reply.started":"2025-04-22T15:49:18.606892Z","shell.execute_reply":"2025-04-22T15:49:18.610516Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"TRAIN_FILENAMES = glob.glob('/kaggle/input/islr-5fold/*.tfrecords')","metadata":{"id":"QVDc3vkwqAFQ","outputId":"b84108e5-9e50-4a46-eca8-6a60ecdde99c","trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:49:21.030612Z","iopub.execute_input":"2025-04-22T15:49:21.030965Z","iopub.status.idle":"2025-04-22T15:49:21.074906Z","shell.execute_reply.started":"2025-04-22T15:49:21.030938Z","shell.execute_reply":"2025-04-22T15:49:21.074031Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ROWS_PER_FRAME = 543\n# MAX_LEN = 384\n# CROP_LEN = MAX_LEN\n# NUM_CLASSES  = 250\n# PAD = -100.\n# NOSE=[\n#     1,2,98,327\n# ]\n# LNOSE = [98]\n# RNOSE = [327]\n# LIP = [ 0, \n#     61, 185, 40, 39, 37, 267, 269, 270, 409,\n#     291, 146, 91, 181, 84, 17, 314, 405, 321, 375,\n#     78, 191, 80, 81, 82, 13, 312, 311, 310, 415,\n#     95, 88, 178, 87, 14, 317, 402, 318, 324, 308,\n# ]\n# LLIP = [84,181,91,146,61,185,40,39,37,87,178,88,95,78,191,80,81,82]\n# RLIP = [314,405,321,375,291,409,270,269,267,317,402,318,324,308,415,310,311,312]\n\n# POSE = [500, 502, 504, 501, 503, 505, 512, 513]\n# LPOSE = [513,505,503,501]\n# RPOSE = [512,504,502,500]\n\n# REYE = [\n#     33, 7, 163, 144, 145, 153, 154, 155, 133,\n#     246, 161, 160, 159, 158, 157, 173,\n# ]\n# LEYE = [\n#     263, 249, 390, 373, 374, 380, 381, 382, 362,\n#     466, 388, 387, 386, 385, 384, 398,\n# ]\n\n# LHAND = np.arange(468, 489).tolist()\n# RHAND = np.arange(522, 543).tolist()\n\n# POINT_LANDMARKS = LIP + LHAND + RHAND + NOSE + REYE + LEYE #+POSE\n\n# NUM_NODES = len(POINT_LANDMARKS)\n# CHANNELS = 6*NUM_NODES\n\n# 定义每帧图像的行数（可能是图像高度或数据点的行数）\nROWS_PER_FRAME = 543\n\n# 定义最大序列长度（可能用于时间序列数据，如视频帧或时间步）\nMAX_LEN = 384\n\n# 定义裁剪长度（与MAX_LEN相同，可能用于数据预处理时的裁剪操作）\nCROP_LEN = MAX_LEN\n\n# 定义分类任务的类别数量（可能是250个不同的动作或姿态类别）\nNUM_CLASSES = 250\n\n# 定义填充值（用于数据填充或特殊标记，如无效值）\nPAD = -100.\n\n# 定义鼻子相关的关键点索引（可能用于面部或姿态估计）\nNOSE = [\n    1,  # 可能是鼻根或鼻梁的某个点\n    2,  # 可能是鼻尖或鼻翼的某个点\n    98, # 中鼻（可能是鼻梁中点）\n    327 # 外鼻（可能是鼻翼外侧）\n]\n\n# 定义左鼻孔的关键点索引（仅包含中鼻）\nLNOSE = [98]\n\n# 定义右鼻孔的关键点索引（仅包含外鼻）\nRNOSE = [327]\n\n# 定义嘴唇的关键点索引（面部轮廓或表情分析）\nLIP = [\n    0,  # 可能是嘴唇中心或某个基准点\n    61, 185, 40, 39, 37, 267, 269, 270, 409,  # 上唇或嘴唇上部关键点\n    291, 146, 91, 181, 84, 17, 314, 405, 321, 375,  # 下唇或嘴唇下部关键点\n    78, 191, 80, 81, 82, 13, 312, 311, 310, 415,  # 嘴唇边缘或其他细节\n    95, 88, 178, 87, 14, 317, 402, 318, 324, 308   # 嘴唇内部或其他特征点\n]\n\n# 定义左嘴唇的关键点索引（从LIP中提取部分点）\nLLIP = [84, 181, 91, 146, 61, 185, 40, 39, 37, 87, 178, 88, 95, 78, 191, 80, 81, 82]\n\n# 定义右嘴唇的关键点索引（从LIP中提取部分点）\nRLIP = [314, 405, 321, 375, 291, 409, 270, 269, 267, 317, 402, 318, 324, 308, 415, 310, 311, 312]\n\n# 定义姿态（可能是身体姿态）相关的关键点索引\nPOSE = [\n    500, 502, 504, 501, 503, 505,  # 可能是身体上半部分的关键点\n    512, 513                     # 可能是身体下半部分的关键点\n]\n\n# 定义左姿态的关键点索引（从POSE中提取部分点）\nLPOSE = [513, 505, 503, 501]  # 可能是左臂或左侧身体的关键点\n\n# 定义右姿态的关键点索引（从POSE中提取部分点）\nRPOSE = [512, 504, 502, 500]  # 可能是右臂或右侧身体的关键点\n\n# 定义右眼的关键点索引（面部表情分析）\nREYE = [\n    33, 7, 163, 144, 145, 153, 154, 155, 133,  # 右眼外部轮廓\n    246, 161, 160, 159, 158, 157, 173        # 右眼内部细节（如瞳孔、眼角）\n]\n\n# 定义左眼的关键点索引（面部表情分析）\nLEYE = [\n    263, 249, 390, 373, 374, 380, 381, 382, 362,  # 左眼外部轮廓\n    466, 388, 387, 386, 385, 384, 398            # 左眼内部细节（如瞳孔、眼角）\n]\n\n# 定义左手的关键点索引（手势识别或姿态估计）\nLHAND = np.arange(468, 489).tolist()  # 手部关键点索引范围（468到488）\n\n# 定义右手的关键点索引（手势识别或姿态估计）\nRHAND = np.arange(522, 543).tolist()  # 手部关键点索引范围（522到542）\n\n# 定义所有需要使用的关键点索引（嘴唇、手部、鼻子、眼睛）\nPOINT_LANDMARKS = LIP + LHAND + RHAND + NOSE + REYE + LEYE  # 可能还有POSE，但被注释掉了\n\n# 计算关键点的总数（用于后续神经网络输入维度计算）\nNUM_NODES = len(POINT_LANDMARKS)  # 关键点数量\n\n# 计算输入通道数（每个关键点可能有6个特征，如x,y坐标 + 其他属性）\nCHANNELS = 6 * NUM_NODES  # 每个关键点的特征维度（如x,y坐标 + 其他4维特征）\n\n\ndef interp1d_(x, target_len, method='random'):\n    length = tf.shape(x)[1]\n    target_len = tf.maximum(1,target_len)\n    if method == 'random':\n        if tf.random.uniform(()) < 0.33:\n            x = tf.image.resize(x, (target_len,tf.shape(x)[1]),'bilinear')\n        else:\n            if tf.random.uniform(()) < 0.5:\n                x = tf.image.resize(x, (target_len,tf.shape(x)[1]),'bicubic')\n            else:\n                x = tf.image.resize(x, (target_len,tf.shape(x)[1]),'nearest')\n    else:\n        x = tf.image.resize(x, (target_len,tf.shape(x)[1]),method)\n    return x\n\ndef tf_nan_mean(x, axis=0, keepdims=False):\n    return tf.reduce_sum(tf.where(tf.math.is_nan(x), tf.zeros_like(x), x), axis=axis, keepdims=keepdims) / tf.reduce_sum(tf.where(tf.math.is_nan(x), tf.zeros_like(x), tf.ones_like(x)), axis=axis, keepdims=keepdims)\n\ndef tf_nan_std(x, center=None, axis=0, keepdims=False):\n    if center is None:\n        center = tf_nan_mean(x, axis=axis,  keepdims=True)\n    d = x - center\n    return tf.math.sqrt(tf_nan_mean(d * d, axis=axis, keepdims=keepdims))\n\nclass Preprocess(tf.keras.layers.Layer):\n    def __init__(self, max_len=MAX_LEN, point_landmarks=POINT_LANDMARKS, **kwargs):\n        super().__init__(**kwargs)\n        self.max_len = max_len\n        self.point_landmarks = point_landmarks\n\n    def call(self, inputs):\n        if tf.rank(inputs) == 3:\n            x = inputs[None,...]\n        else:\n            x = inputs\n        \n        mean = tf_nan_mean(tf.gather(x, [17], axis=2), axis=[1,2], keepdims=True)\n        mean = tf.where(tf.math.is_nan(mean), tf.constant(0.5,x.dtype), mean)\n        x = tf.gather(x, self.point_landmarks, axis=2) #N,T,P,C\n        std = tf_nan_std(x, center=mean, axis=[1,2], keepdims=True)\n        \n        x = (x - mean)/std\n\n        if self.max_len is not None:\n            x = x[:,:self.max_len]\n        length = tf.shape(x)[1]\n        x = x[...,:2]\n\n        dx = tf.cond(tf.shape(x)[1]>1,lambda:tf.pad(x[:,1:] - x[:,:-1], [[0,0],[0,1],[0,0],[0,0]]),lambda:tf.zeros_like(x))\n\n        dx2 = tf.cond(tf.shape(x)[1]>2,lambda:tf.pad(x[:,2:] - x[:,:-2], [[0,0],[0,2],[0,0],[0,0]]),lambda:tf.zeros_like(x))\n\n        x = tf.concat([\n            tf.reshape(x, (-1,length,2*len(self.point_landmarks))),\n            tf.reshape(dx, (-1,length,2*len(self.point_landmarks))),\n            tf.reshape(dx2, (-1,length,2*len(self.point_landmarks))),\n        ], axis = -1)\n        \n        x = tf.where(tf.math.is_nan(x),tf.constant(0.,x.dtype),x)\n        \n        return x","metadata":{"id":"6xyloTyiqAFQ","outputId":"dab46b40-8d61-4eb7-a7a5-4d5e10ab8538","trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:49:52.735376Z","iopub.execute_input":"2025-04-22T15:49:52.735751Z","iopub.status.idle":"2025-04-22T15:49:52.761810Z","shell.execute_reply.started":"2025-04-22T15:49:52.735726Z","shell.execute_reply":"2025-04-22T15:49:52.761068Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 功能：解析TFRecord样本。\n# 流程：\n# 解析二进制数据中的coordinates（关键点坐标）和sign（标签）。\n# coordinates被解码为形状(T, 543, 3)的浮点张量，表示时间步×关键点×坐标（x,y,z）。\n# sign转换为int64标签。\n\ndef decode_tfrec(record_bytes):\n    features = tf.io.parse_single_example(record_bytes, {\n        'coordinates': tf.io.FixedLenFeature([], tf.string),\n        'sign': tf.io.FixedLenFeature([], tf.int64),\n    })\n    out = {}\n    out['coordinates']  = tf.reshape(tf.io.decode_raw(features['coordinates'], tf.float32), (-1,ROWS_PER_FRAME,3))\n    out['sign'] = features['sign']\n    return out\n\n# 功能：过滤存在NaN的无效帧。\n# 逻辑：\n# 检查ref_point（关键参考点）是否存在全NaN。\n# 通过掩码保留至少有一个有效参考点的帧。\n\ndef filter_nans_tf(x, ref_point=POINT_LANDMARKS):\n    mask = tf.math.logical_not(tf.reduce_all(tf.math.is_nan(tf.gather(x,ref_point,axis=1)), axis=[-2,-1]))\n    x = tf.boolean_mask(x, mask, axis=0)\n    return x\n\n\n\n# preprocess(x, augment)\n# 流程：\n\n# 过滤NaN帧。\n\n# 可选增强（augment=True时调用augment_fn）。\n\n# 应用Preprocess层标准化并提取运动特征。\n\n# 标签转为One-hot编码。\ndef preprocess(x, augment=False, max_len=MAX_LEN):\n    coord = x['coordinates']\n    coord = filter_nans_tf(coord)\n    if augment:\n        coord = augment_fn(coord, max_len=max_len)\n    coord = tf.ensure_shape(coord, (None,ROWS_PER_FRAME,3))\n    \n    return tf.cast(Preprocess(max_len=max_len)(coord)[0],tf.float32), tf.one_hot(x['sign'], NUM_CLASSES)\n\n\n\n# 功能：左右镜像翻转。\n# 关键操作：\n# 水平翻转x坐标：x = 1 - x。\n# 对称关键点交换：\n# 左右手、嘴唇、姿势、眼睛、鼻子等关键点索引互换（如左手变右手）。\n# 保持空间合理性，避免模型对左右方向过拟合。\n\ndef flip_lr(x):\n    x,y,z = tf.unstack(x, axis=-1)\n    x = 1-x\n    new_x = tf.stack([x,y,z], -1)\n    new_x = tf.transpose(new_x, [1,0,2])\n    lhand = tf.gather(new_x, LHAND, axis=0)\n    rhand = tf.gather(new_x, RHAND, axis=0)\n    new_x = tf.tensor_scatter_nd_update(new_x, tf.constant(LHAND)[...,None], rhand)\n    new_x = tf.tensor_scatter_nd_update(new_x, tf.constant(RHAND)[...,None], lhand)\n    llip = tf.gather(new_x, LLIP, axis=0)\n    rlip = tf.gather(new_x, RLIP, axis=0)\n    new_x = tf.tensor_scatter_nd_update(new_x, tf.constant(LLIP)[...,None], rlip)\n    new_x = tf.tensor_scatter_nd_update(new_x, tf.constant(RLIP)[...,None], llip)\n    lpose = tf.gather(new_x, LPOSE, axis=0)\n    rpose = tf.gather(new_x, RPOSE, axis=0)\n    new_x = tf.tensor_scatter_nd_update(new_x, tf.constant(LPOSE)[...,None], rpose)\n    new_x = tf.tensor_scatter_nd_update(new_x, tf.constant(RPOSE)[...,None], lpose)\n    leye = tf.gather(new_x, LEYE, axis=0)\n    reye = tf.gather(new_x, REYE, axis=0)\n    new_x = tf.tensor_scatter_nd_update(new_x, tf.constant(LEYE)[...,None], reye)\n    new_x = tf.tensor_scatter_nd_update(new_x, tf.constant(REYE)[...,None], leye)\n    lnose = tf.gather(new_x, LNOSE, axis=0)\n    rnose = tf.gather(new_x, RNOSE, axis=0)\n    new_x = tf.tensor_scatter_nd_update(new_x, tf.constant(LNOSE)[...,None], rnose)\n    new_x = tf.tensor_scatter_nd_update(new_x, tf.constant(RNOSE)[...,None], lnose)\n    new_x = tf.transpose(new_x, [1,0,2])\n    return new_x\n\n\n# 功能：时间维度重采样。\n# 实现\n# 随机缩放帧率（如0.5倍慢放或1.5倍快放）。\n# 使用插值调整序列长度（interp1d_函数）。\n\ndef resample(x, rate=(0.8,1.2)):\n    rate = tf.random.uniform((), rate[0], rate[1])\n    length = tf.shape(x)[0]\n    new_size = tf.cast(rate*tf.cast(length,tf.float32), tf.int32)\n    new_x = interp1d_(x, new_size)\n    return new_x\n\n\n# 功能：空间仿射变换增强。\n# 变换类型：\n# 缩放：随机缩放关键点坐标。\n# 剪切：模拟非正面视角的变形。\n# 旋转：绕中心点随机旋转。\n# 平移：整体偏移位置。\n# 目的：增强模型对视角、距离变化的鲁棒性。\ndef spatial_random_affine(xyz,\n    scale  = (0.8,1.2),\n    shear = (-0.15,0.15),\n    shift  = (-0.1,0.1),\n    degree = (-30,30),\n):\n    center = tf.constant([0.5,0.5])\n    if scale is not None:\n        scale = tf.random.uniform((),*scale)\n        xyz = scale*xyz\n\n    if shear is not None:\n        xy = xyz[...,:2]\n        z = xyz[...,2:]\n        shear_x = shear_y = tf.random.uniform((),*shear)\n        if tf.random.uniform(()) < 0.5:\n            shear_x = 0.\n        else:\n            shear_y = 0.\n        shear_mat = tf.identity([\n            [1.,shear_x],\n            [shear_y,1.]\n        ])\n        xy = xy @ shear_mat\n        center = center + [shear_y, shear_x]\n        xyz = tf.concat([xy,z], axis=-1)\n\n    if degree is not None:\n        xy = xyz[...,:2]\n        z = xyz[...,2:]\n        xy -= center\n        degree = tf.random.uniform((),*degree)\n        radian = degree/180*np.pi\n        c = tf.math.cos(radian)\n        s = tf.math.sin(radian)\n        rotate_mat = tf.identity([\n            [c,s],\n            [-s, c],\n        ])\n        xy = xy @ rotate_mat\n        xy = xy + center\n        xyz = tf.concat([xy,z], axis=-1)\n\n    if shift is not None:\n        shift = tf.random.uniform((),*shift)\n        xyz = xyz + shift\n\n    return xyz\n\n\n# 功能：随机截取时间片段。\n# 作用：统一输入长度，增加时间维度多样性。\ndef temporal_crop(x, length=MAX_LEN):\n    l = tf.shape(x)[0]\n    offset = tf.random.uniform((), 0, tf.clip_by_value(l-length,1,length), dtype=tf.int32)\n    x = x[offset:offset+length]\n    return x\n\n\n# 功能：随机遮蔽连续时间段（置NaN）。\n# 效果：模拟视频中的遮挡或丢帧，迫使模型关注局部时序特征。\ndef temporal_mask(x, size=(0.2,0.4), mask_value=float('nan')):\n    l = tf.shape(x)[0]\n    mask_size = tf.random.uniform((), *size)\n    mask_size = tf.cast(tf.cast(l, tf.float32) * mask_size, tf.int32)\n    mask_offset = tf.random.uniform((), 0, tf.clip_by_value(l-mask_size,1,l), dtype=tf.int32)\n    x = tf.tensor_scatter_nd_update(x,tf.range(mask_offset, mask_offset+mask_size)[...,None],tf.fill([mask_size,543,3],mask_value))\n    return x\n\n\n# 功能：随机遮蔽空间区域（矩形区域置NaN）。\n# 效果：模拟部分关键点检测失败，增强模型对局部特征的关注。\ndef spatial_mask(x, size=(0.2,0.4), mask_value=float('nan')):\n    mask_offset_y = tf.random.uniform(())\n    mask_offset_x = tf.random.uniform(())\n    mask_size = tf.random.uniform((), *size)\n    mask_x = (mask_offset_x<x[...,0]) & (x[...,0] < mask_offset_x + mask_size)\n    mask_y = (mask_offset_y<x[...,1]) & (x[...,1] < mask_offset_y + mask_size)\n    mask = mask_x & mask_y\n    x = tf.where(mask[...,None], mask_value, x)\n    return x\n\n\n# 功能：组合增强策略。\n# 随机应用：\n# 重采样（80%概率）。\n# 左右翻转（50%概率）。\n# 时间裁剪（若指定max_len）。\n# 空间仿射变换（75%概率）。\n# 时间遮蔽（50%概率）。\n# 空间遮蔽（50%概率）。\ndef augment_fn(x, always=False, max_len=None):\n    if tf.random.uniform(())<0.8 or always:\n        x = resample(x, (0.5,1.5))\n    if tf.random.uniform(())<0.5 or always:\n        x = flip_lr(x)\n    if max_len is not None:\n        x = temporal_crop(x, max_len)\n    if tf.random.uniform(())<0.75 or always:\n        x = spatial_random_affine(x)\n    if tf.random.uniform(())<0.5 or always:\n        x = temporal_mask(x)\n    if tf.random.uniform(())<0.5 or always:\n        x = spatial_mask(x)\n    return x\n\n\n\n# 功能：构建TensorFlow输入流水线。\n# 步骤：\n# 读取TFRecord文件：并行解码（num_parallel_reads=AUTOTUNE）。\n# 映射预处理：调用decode_tfrec和preprocess。\n# 重复与乱序：根据参数控制是否重复和打乱顺序。\n# 批处理：动态填充为统一长度（max_len），填充值PAD=-100。\n# 性能优化：预取数据（prefetch）确保GPU不空闲。\ndef get_tfrec_dataset(tfrecords, batch_size=64, max_len=64, drop_remainder=False, augment=False, shuffle=False, repeat=False):\n    # Initialize dataset with TFRecords\n    ds = tf.data.TFRecordDataset(tfrecords, num_parallel_reads=tf.data.AUTOTUNE, compression_type='GZIP')\n    ds = ds.map(decode_tfrec, tf.data.AUTOTUNE)\n    ds = ds.map(lambda x: preprocess(x, augment=augment, max_len=max_len), tf.data.AUTOTUNE)\n\n    if repeat: \n        ds = ds.repeat()\n        \n    if shuffle:\n        ds = ds.shuffle(shuffle)\n        options = tf.data.Options()\n        options.experimental_deterministic = (False)\n        ds = ds.with_options(options)\n    \n    if batch_size:\n        ds = ds.padded_batch(batch_size, padding_values=PAD, padded_shapes=([max_len,CHANNELS],[NUM_CLASSES]), drop_remainder=drop_remainder)\n\n    ds = ds.prefetch(tf.data.AUTOTUNE)\n        \n\n    return ds","metadata":{"id":"r17ZnZaGqAFQ","trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:49:59.414175Z","iopub.execute_input":"2025-04-22T15:49:59.415100Z","iopub.status.idle":"2025-04-22T15:49:59.452529Z","shell.execute_reply.started":"2025-04-22T15:49:59.415064Z","shell.execute_reply":"2025-04-22T15:49:59.451679Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\n1. ECA (Efficient Channel Attention) 通道注意力机制\n功能：自适应学习通道权重，增强重要特征。\n\n实现：\n\n全局平均池化（GlobalAveragePooling1D）获取通道级统计量。\n\n1D卷积（Conv1D）学习通道间关系，卷积核大小可调（默认5）。\n\nSigmoid激活生成0-1的通道权重。\n\n与原始输入逐通道相乘，实现特征校准。\n\"\"\"\nclass ECA(tf.keras.layers.Layer):\n    def __init__(self, kernel_size=5, **kwargs):\n        super().__init__(**kwargs)\n        self.supports_masking = True\n        self.kernel_size = kernel_size\n        self.conv = tf.keras.layers.Conv1D(1, kernel_size=kernel_size, strides=1, padding=\"same\", use_bias=False)\n\n    def call(self, inputs, mask=None):\n        nn = tf.keras.layers.GlobalAveragePooling1D()(inputs, mask=mask)\n        nn = tf.expand_dims(nn, -1)\n        nn = self.conv(nn)\n        nn = tf.squeeze(nn, -1)\n        nn = tf.nn.sigmoid(nn)\n        nn = nn[:,None,:]\n        return inputs * nn\n\n\n\"\"\"\n2. LateDropout 延迟丢弃层\n功能：在指定训练步数后启用Dropout，兼顾训练稳定性与正则化。\n\n机制：\n\n内置计数器_train_counter记录训练步数。\n\n当_train_counter >= start_step时应用Dropout。\n\n适用于逐步引入正则化，避免训练初期过拟合。\n\"\"\"\nclass LateDropout(tf.keras.layers.Layer):\n    def __init__(self, rate, noise_shape=None, start_step=0, **kwargs):\n        super().__init__(**kwargs)\n        self.supports_masking = True\n        self.rate = rate\n        self.start_step = start_step\n        self.dropout = tf.keras.layers.Dropout(rate, noise_shape=noise_shape)\n      \n    def build(self, input_shape):\n        super().build(input_shape)\n        agg = tf.VariableAggregation.ONLY_FIRST_REPLICA\n        self._train_counter = tf.Variable(0, dtype=\"int64\", aggregation=agg, trainable=False)\n\n    def call(self, inputs, training=False):\n        x = tf.cond(self._train_counter < self.start_step, lambda:inputs, lambda:self.dropout(inputs, training=training))\n        if training:\n            self._train_counter.assign_add(1)\n        return x\n\n\n\"\"\"\n3. CausalDWConv1D 因果深度可分离卷积\n因果性：通过左填充确保卷积不依赖未来信息。\n\n结构：\n\nZeroPadding1D：填充左侧(dilation_rate*(kernel_size-1), 0)。\n\nDepthwiseConv1D：深度卷积提取时序特征。\n\"\"\"\nclass CausalDWConv1D(tf.keras.layers.Layer):\n    def __init__(self, \n        kernel_size=17,\n        dilation_rate=1,\n        use_bias=False,\n        depthwise_initializer='glorot_uniform',\n        name='', **kwargs):\n        super().__init__(name=name,**kwargs)\n        self.causal_pad = tf.keras.layers.ZeroPadding1D((dilation_rate*(kernel_size-1),0),name=name + '_pad')\n        self.dw_conv = tf.keras.layers.DepthwiseConv1D(\n                            kernel_size,\n                            strides=1,\n                            dilation_rate=dilation_rate,\n                            padding='valid',\n                            use_bias=use_bias,\n                            depthwise_initializer=depthwise_initializer,\n                            name=name + '_dwconv')\n        self.supports_masking = True\n        \n    def call(self, inputs):\n        x = self.causal_pad(inputs)\n        x = self.dw_conv(x)\n        return x\n\n\n\"\"\"\n4. Conv1DBlock 高效卷积块\n结构：\n\n通道扩展：Dense层扩展通道（expand_ratio=2）。\n\n因果深度卷积：CausalDWConv1D提取局部时序特征。\n\nECA注意力：增强关键通道。\n\n通道压缩：Dense层恢复原始通道数。\n\n残差连接：输入输出通道相同时添加跳跃连接。\n\"\"\"\ndef Conv1DBlock(channel_size,\n          kernel_size,\n          dilation_rate=1,\n          drop_rate=0.0,\n          expand_ratio=2,\n          se_ratio=0.25,\n          activation='swish',\n          name=None):\n    '''\n    efficient conv1d block, @hoyso48\n    '''\n    if name is None:\n        name = str(tf.keras.backend.get_uid(\"mbblock\"))\n    # Expansion phase\n    def apply(inputs):\n        channels_in = tf.keras.backend.int_shape(inputs)[-1]\n        channels_expand = channels_in * expand_ratio\n\n        skip = inputs\n\n        x = tf.keras.layers.Dense(\n            channels_expand,\n            use_bias=True,\n            activation=activation,\n            name=name + '_expand_conv')(inputs)\n\n        # Depthwise Convolution\n        x = CausalDWConv1D(kernel_size,\n            dilation_rate=dilation_rate,\n            use_bias=False,\n            name=name + '_dwconv')(x)\n\n        x = tf.keras.layers.BatchNormalization(momentum=0.95, name=name + '_bn')(x)\n\n        x  = ECA()(x)\n\n        x = tf.keras.layers.Dense(\n            channel_size,\n            use_bias=True,\n            name=name + '_project_conv')(x)\n\n        if drop_rate > 0:\n            x = tf.keras.layers.Dropout(drop_rate, noise_shape=(None,1,1), name=name + '_drop')(x)\n\n        if (channels_in == channel_size):\n            x = tf.keras.layers.add([x, skip], name=name + '_add')\n        return x\n\n    return apply","metadata":{"id":"zaIHc89AqAFQ","trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:50:04.586060Z","iopub.execute_input":"2025-04-22T15:50:04.586430Z","iopub.status.idle":"2025-04-22T15:50:04.604587Z","shell.execute_reply.started":"2025-04-22T15:50:04.586403Z","shell.execute_reply":"2025-04-22T15:50:04.603697Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\n5. MultiHeadSelfAttention 多头自注意力\n功能：捕捉长距离时序依赖。\n\n实现：\n\n生成Q、K、V矩阵并分割为多头。\n\n缩放点积注意力计算。\n\n合并多头输出并通过投影层。\n\"\"\"\nclass MultiHeadSelfAttention(tf.keras.layers.Layer):\n    def __init__(self, dim=256, num_heads=4, dropout=0, **kwargs):\n        super().__init__(**kwargs)\n        self.dim = dim\n        self.scale = self.dim ** -0.5\n        self.num_heads = num_heads\n        self.qkv = tf.keras.layers.Dense(3 * dim, use_bias=False)\n        self.drop1 = tf.keras.layers.Dropout(dropout)\n        self.proj = tf.keras.layers.Dense(dim, use_bias=False)\n        self.supports_masking = True\n\n    def call(self, inputs, mask=None):\n        qkv = self.qkv(inputs)\n        qkv = tf.keras.layers.Permute((2, 1, 3))(tf.keras.layers.Reshape((-1, self.num_heads, self.dim * 3 // self.num_heads))(qkv))\n        q, k, v = tf.split(qkv, [self.dim // self.num_heads] * 3, axis=-1)\n\n        attn = tf.matmul(q, k, transpose_b=True) * self.scale\n\n        if mask is not None:\n            mask = mask[:, None, None, :]\n\n        attn = tf.keras.layers.Softmax(axis=-1)(attn, mask=mask)\n        attn = self.drop1(attn)\n\n        x = attn @ v\n        x = tf.keras.layers.Reshape((-1, self.dim))(tf.keras.layers.Permute((2, 1, 3))(x))\n        x = self.proj(x)\n        return x\n\n\"\"\"\n6. TransformerBlock Transformer模块\n结构：\n\n自注意力子层：MultiHeadSelfAttention + Dropout + 残差。\n\n前馈子层：Dense扩展（expand=4）→ Dense压缩 → Dropout + 残差。\n\nBatchNorm替代LayerNorm：与标准Transformer不同。\n\"\"\"\ndef TransformerBlock(dim=256, num_heads=4, expand=4, attn_dropout=0.2, drop_rate=0.2, activation='swish'):\n    def apply(inputs):\n        x = inputs\n        x = tf.keras.layers.BatchNormalization(momentum=0.95)(x)\n        x = MultiHeadSelfAttention(dim=dim,num_heads=num_heads,dropout=attn_dropout)(x)\n        x = tf.keras.layers.Dropout(drop_rate, noise_shape=(None,1,1))(x)\n        x = tf.keras.layers.Add()([inputs, x])\n        attn_out = x\n\n        x = tf.keras.layers.BatchNormalization(momentum=0.95)(x)\n        x = tf.keras.layers.Dense(dim*expand, use_bias=False, activation=activation)(x)\n        x = tf.keras.layers.Dense(dim, use_bias=False)(x)\n        x = tf.keras.layers.Dropout(drop_rate, noise_shape=(None,1,1))(x)\n        x = tf.keras.layers.Add()([attn_out, x])\n        return x\n    return apply","metadata":{"id":"8hPmJX0YqAFR","trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:50:08.404145Z","iopub.execute_input":"2025-04-22T15:50:08.404533Z","iopub.status.idle":"2025-04-22T15:50:08.417524Z","shell.execute_reply.started":"2025-04-22T15:50:08.404503Z","shell.execute_reply":"2025-04-22T15:50:08.416736Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 新增\nclass WindowedAttention(tf.keras.layers.Layer):\n    def __init__(self, dim=256, num_heads=4, window_size=16, shift=False, dropout=0.2):\n        super().__init__()\n        self.dim = dim\n        self.num_heads = num_heads\n        self.window_size = window_size\n        self.shift = shift\n        \n        # 相对位置编码矩阵 (2*window_size-1)\n        self.rel_pos_bias = self.add_weight(\n            name=\"rel_pos_bias\",\n            shape=(2 * window_size - 1, num_heads),\n            initializer=\"zeros\",\n        )\n        \n        self.qkv = tf.keras.layers.Dense(3*dim)\n        self.proj = tf.keras.layers.Dense(dim)\n        self.dropout = tf.keras.layers.Dropout(dropout)\n\n    def get_relative_positions(self, length):\n        # 生成相对位置索引矩阵\n        range_vec = tf.range(length)\n        distance_mat = range_vec[:, None] - range_vec[None, :]\n        return distance_mat + self.window_size - 1  # 偏移至非负数\n\n    def window_partition(self, x):\n        # x: (B, T, C)\n        B, T, C = tf.shape(x)[0], tf.shape(x)[1], tf.shape(x)[2]\n        \n        # 移位操作\n        if self.shift:\n            x = tf.roll(x, shift=-self.window_size//2, axis=1)\n        \n        # 填充与分窗\n        pad_len = (self.window_size - T % self.window_size) % self.window_size\n        x = tf.pad(x, [[0,0], [0,pad_len], [0,0]])\n        x = tf.reshape(x, [B, -1, self.window_size, C])  # (B, num_win, win_size, C)\n        return x, pad_len\n\n    def call(self, x, mask=None):\n        B, T, C = tf.shape(x)[0], tf.shape(x)[1], tf.shape(x)[2]\n        \n        # 分窗处理\n        x_windows, pad_len = self.window_partition(x)  # (B, nW, win, C)\n        x_windows = tf.reshape(x_windows, [-1, self.window_size, C])  # (B*nW, win, C)\n        \n        # 生成QKV\n        qkv = self.qkv(x_windows)\n        qkv = tf.reshape(qkv, [-1, self.window_size, 3, self.num_heads, C//self.num_heads])\n        q, k, v = tf.unstack(qkv, axis=2)  # 各 (B*nW, win, nH, C/nH)\n        \n        # 相对位置编码\n        rel_pos = self.get_relative_positions(self.window_size)\n        rel_bias = tf.gather(self.rel_pos_bias, rel_pos)  # (win, win, nH)\n        rel_bias = tf.transpose(rel_bias, [2,0,1])  # (nH, win, win)\n        \n        # 注意力计算\n        attn = tf.einsum('bqhd,bkhd->bhqk', q, k)  # (B*nW, nH, win, win)\n        attn = attn + rel_bias[None,...]  # 加入位置偏置\n        attn = attn / tf.sqrt(tf.cast(C//self.num_heads, tf.float32))\n        \n        if mask is not None:\n            # 继承原有mask机制\n            mask = tf.reshape(mask, [B, -1, self.window_size])\n            mask = tf.repeat(mask[:, None, :, :], self.num_heads, axis=1)\n            attn = tf.where(mask, attn, -1e9)\n        \n        attn = tf.nn.softmax(attn, axis=-1)\n        attn = self.dropout(attn)\n        \n        # 聚合Value\n        out = tf.einsum('bhqk,bkhd->bqhd', attn, v)\n        out = tf.reshape(out, [B, -1, self.dim])  # (B, T_pad, C)\n        \n        # 移除填充\n        if pad_len > 0:\n            out = out[:, :T, :]\n        \n        return self.proj(out)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:50:11.815751Z","iopub.execute_input":"2025-04-22T15:50:11.816469Z","iopub.status.idle":"2025-04-22T15:50:11.831856Z","shell.execute_reply.started":"2025-04-22T15:50:11.816432Z","shell.execute_reply":"2025-04-22T15:50:11.831042Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 新增\ndef SwinTransformerBlock(dim=256, num_heads=4, window_size=16, shift=False, expand=4):\n    def apply(inputs):\n        x = inputs\n        \n        # 窗口注意力\n        x_norm = tf.keras.layers.LayerNormalization()(x)\n        attn = WindowedAttention(dim, num_heads, window_size, shift)(x_norm)\n        x = x + attn  # 残差连接\n        \n        # 前馈网络\n        x_norm = tf.keras.layers.LayerNormalization()(x)\n        ffn = tf.keras.Sequential([\n            tf.keras.layers.Dense(dim * expand, activation='gelu'),\n            tf.keras.layers.Dense(dim)\n        ])(x_norm)\n        x = x + ffn\n        \n        return x\n    return apply","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:50:15.720793Z","iopub.execute_input":"2025-04-22T15:50:15.721492Z","iopub.status.idle":"2025-04-22T15:50:15.727345Z","shell.execute_reply.started":"2025-04-22T15:50:15.721455Z","shell.execute_reply":"2025-04-22T15:50:15.726579Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nInputLayer: (None, 64, 6*NUM_NODES)\n    │\n    ▼\nMasking (忽略PAD值)\n    │\n    ▼\nDense (dim=192) + BatchNorm  # 初始特征映射\n    │\n    ▼\nConv1DBlock → Conv1DBlock → Conv1DBlock → TransformerBlock  # 阶段1\n    │\n    ▼\nConv1DBlock → Conv1DBlock → Conv1DBlock → TransformerBlock  # 阶段2\n    │\n    ▼（仅当dim=384时）\nConv1DBlock ×3 → TransformerBlock ×2（扩展阶段）\n    │\n    ▼\nDense (dim*2) → GlobalAvgPool → LateDropout → Dense(250)  # 分类\n\"\"\"\ndef get_model(max_len=64, dropout_step=0, dim=192):\n    inp = tf.keras.Input((max_len,CHANNELS))\n    x = tf.keras.layers.Masking(mask_value=PAD,input_shape=(max_len,CHANNELS))(inp)\n    ksize = 17\n    x = tf.keras.layers.Dense(dim, use_bias=False,name='stem_conv')(x)\n    x = tf.keras.layers.BatchNormalization(momentum=0.95,name='stem_bn')(x)\n\n    x = Conv1DBlock(dim,ksize,drop_rate=0.2)(x)\n    x = Conv1DBlock(dim,ksize,drop_rate=0.2)(x)\n    x = Conv1DBlock(dim,ksize,drop_rate=0.2)(x)\n    x = TransformerBlock(dim,expand=2)(x)\n\n    x = Conv1DBlock(dim,ksize,drop_rate=0.2)(x)\n    x = Conv1DBlock(dim,ksize,drop_rate=0.2)(x)\n    x = Conv1DBlock(dim,ksize,drop_rate=0.2)(x)\n    x = TransformerBlock(dim,expand=2)(x)\n\n    if dim == 384: #for the 4x sized model\n        x = Conv1DBlock(dim,ksize,drop_rate=0.2)(x)\n        x = Conv1DBlock(dim,ksize,drop_rate=0.2)(x)\n        x = Conv1DBlock(dim,ksize,drop_rate=0.2)(x)\n        x = TransformerBlock(dim,expand=2)(x)\n\n        x = Conv1DBlock(dim,ksize,drop_rate=0.2)(x)\n        x = Conv1DBlock(dim,ksize,drop_rate=0.2)(x)\n        x = Conv1DBlock(dim,ksize,drop_rate=0.2)(x)\n        x = TransformerBlock(dim,expand=2)(x)\n\n    x = tf.keras.layers.Dense(dim*2,activation=None,name='top_conv')(x)\n    x = tf.keras.layers.GlobalAveragePooling1D()(x)\n    x = LateDropout(0.8, start_step=dropout_step)(x)\n    x = tf.keras.layers.Dense(NUM_CLASSES,name='classifier')(x)\n    return tf.keras.Model(inp, x)","metadata":{"id":"KIooIcnSqAFR","outputId":"4d6a1f52-1630-45d6-a33c-441aab02d1b2","trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:50:20.882539Z","iopub.execute_input":"2025-04-22T15:50:20.883396Z","iopub.status.idle":"2025-04-22T15:50:20.893698Z","shell.execute_reply.started":"2025-04-22T15:50:20.883331Z","shell.execute_reply":"2025-04-22T15:50:20.892895Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 新模型\ndef get_swin_model(max_len=384, dim=192):\n    inp = tf.keras.Input((max_len, CHANNELS))\n    x = tf.keras.layers.Masking(mask_value=PAD)(inp)\n    \n    # Stem层（保留原卷积特征提取）\n    x = tf.keras.layers.Dense(dim, use_bias=False, name='stem_conv')(x)\n    x = tf.keras.layers.BatchNormalization(momentum=0.95, name='stem_bn')(x)\n    \n    # Stage 1: 小窗口细粒度特征\n    x = Conv1DBlock(dim, 17, drop_rate=0.2)(x)  \n    x = SwinTransformerBlock(dim, window_size=32, shift=False)(x)\n    x = Conv1DBlock(dim, 17, drop_rate=0.2)(x)\n    x = SwinTransformerBlock(dim, window_size=32, shift=True)(x)  # 移位窗口\n    \n    # Stage 2: 中等窗口 + 下采样\n    x = tf.keras.layers.AvgPool1D(2)(x)  # T=192\n    x = Conv1DBlock(dim, 17, drop_rate=0.2)(x)\n    x = SwinTransformerBlock(dim, window_size=24, shift=False)(x)\n    x = Conv1DBlock(dim, 17, drop_rate=0.2)(x)\n    x = SwinTransformerBlock(dim, window_size=24, shift=True)(x)\n    \n    # Stage 3: 大窗口全局上下文\n    x = tf.keras.layers.AvgPool1D(2)(x)  # T=96\n    x = Conv1DBlock(dim, 17, drop_rate=0.2)(x)\n    x = SwinTransformerBlock(dim, window_size=48, shift=False)(x)\n    x = Conv1DBlock(dim, 17, drop_rate=0.2)(x)\n    x = SwinTransformerBlock(dim, window_size=48, shift=True)(x)\n    \n    # 分类头（保持原结构）\n    x = tf.keras.layers.Dense(dim*2, activation=None, name='top_conv')(x)\n    x = tf.keras.layers.GlobalAveragePooling1D()(x)\n    x = LateDropout(0.8, start_step=1000)(x)\n    x = tf.keras.layers.Dense(NUM_CLASSES, name='classifier')(x)\n    \n    return tf.keras.Model(inp, x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:50:24.771407Z","iopub.execute_input":"2025-04-22T15:50:24.772207Z","iopub.status.idle":"2025-04-22T15:50:24.781936Z","shell.execute_reply.started":"2025-04-22T15:50:24.772177Z","shell.execute_reply":"2025-04-22T15:50:24.781129Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = get_model()\nmodel.summary()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:50:55.901676Z","iopub.execute_input":"2025-04-22T15:50:55.902444Z","iopub.status.idle":"2025-04-22T15:50:57.330141Z","shell.execute_reply.started":"2025-04-22T15:50:55.902412Z","shell.execute_reply":"2025-04-22T15:50:57.329232Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"swin_model = get_swin_model()\nswin_model.summary()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:51:05.815860Z","iopub.execute_input":"2025-04-22T15:51:05.816592Z","iopub.status.idle":"2025-04-22T15:51:08.278719Z","shell.execute_reply.started":"2025-04-22T15:51:05.816560Z","shell.execute_reply":"2025-04-22T15:51:08.277844Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CFG:\n    n_splits = 5\n    save_output = True\n    output_dir = '/kaggle/working'\n    \n    seed = 42\n    verbose = 2 #0) silent 1) progress bar 2) one line per epoch\n    \n    max_len = 384\n    replicas = 8\n    lr = 5e-4 * replicas\n    weight_decay = 0.1\n    lr_min = 1e-6\n    epoch = 300\n    warmup = 0\n    batch_size = 64 * replicas\n    snapshot_epochs = []\n    swa_epochs = [] #list(range(epoch//2,epoch+1))\n    \n    fp16 = True\n    fgm = False\n    awp = True\n    awp_lambda = 0.2\n    awp_start_epoch = 15\n    dropout_start_epoch = 15\n    resume = 0\n    decay_type = 'cosine'\n    dim = 192\n    comment = f'islr-fp16-192-8-seed{seed}'","metadata":{"id":"WYUI6mAlqAFR","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_fold(CFG, fold, train_files, valid_files=None, strategy=STRATEGY, summary=True):\n    seed_everything(CFG.seed)\n    tf.keras.backend.clear_session()\n    gc.collect()\n    tf.config.optimizer.set_jit(True)\n        \n    if CFG.fp16:\n        try:\n            policy = mixed_precision.Policy('mixed_bfloat16')\n            mixed_precision.set_global_policy(policy)\n        except:\n            policy = mixed_precision.Policy('mixed_float16')\n            mixed_precision.set_global_policy(policy)\n    else:\n        policy = mixed_precision.Policy('float32')\n        mixed_precision.set_global_policy(policy)\n\n    if fold != 'all':\n        train_ds = get_tfrec_dataset(train_files, batch_size=CFG.batch_size, max_len=CFG.max_len, drop_remainder=True, augment=True, repeat=True, shuffle=32768)\n        valid_ds = get_tfrec_dataset(valid_files, batch_size=CFG.batch_size, max_len=CFG.max_len, drop_remainder=False, repeat=False, shuffle=False)\n    else:\n        train_ds = get_tfrec_dataset(train_files, batch_size=CFG.batch_size, max_len=CFG.max_len, drop_remainder=False, augment=True, repeat=True, shuffle=32768)\n        valid_ds = None\n        valid_files = []\n    \n    num_train = count_data_items(train_files)\n    num_valid = count_data_items(valid_files)\n    steps_per_epoch = num_train//CFG.batch_size\n    with strategy.scope():\n        dropout_step = CFG.dropout_start_epoch * steps_per_epoch\n        model = get_model(max_len=CFG.max_len, dropout_step=dropout_step, dim=CFG.dim)\n\n        schedule = OneCycleLR(CFG.lr, CFG.epoch, warmup_epochs=CFG.epoch*CFG.warmup, steps_per_epoch=steps_per_epoch, resume_epoch=CFG.resume, decay_epochs=CFG.epoch, lr_min=CFG.lr_min, decay_type=CFG.decay_type, warmup_type='linear')\n        decay_schedule = OneCycleLR(CFG.lr*CFG.weight_decay, CFG.epoch, warmup_epochs=CFG.epoch*CFG.warmup, steps_per_epoch=steps_per_epoch, resume_epoch=CFG.resume, decay_epochs=CFG.epoch, lr_min=CFG.lr_min*CFG.weight_decay, decay_type=CFG.decay_type, warmup_type='linear')\n                \n        awp_step = CFG.awp_start_epoch * steps_per_epoch\n        if CFG.fgm:\n            model = FGM(model.input, model.output, delta=CFG.awp_lambda, eps=0., start_step=awp_step)\n        elif CFG.awp:\n            model = AWP(model.input, model.output, delta=CFG.awp_lambda, eps=0., start_step=awp_step)\n\n        opt = tfa.optimizers.RectifiedAdam(learning_rate=schedule, weight_decay=decay_schedule, sma_threshold=4, clipvalue=1.)\n        opt = tfa.optimizers.Lookahead(opt,sync_period=5)\n\n        model.compile(\n            optimizer=opt,\n            loss=[tf.keras.losses.CategoricalCrossentropy(from_logits=True, label_smoothing=0.1)], #[tf.keras.losses.CategoricalCrossentropy(from_logits=True)],\n            metrics=[\n                [\n                tf.keras.metrics.CategoricalAccuracy(),\n                ],\n            ],\n            steps_per_execution=steps_per_epoch,\n        )\n    \n    if summary:\n        print()\n        model.summary()\n        print()\n        print(train_ds, valid_ds)\n        print()\n        schedule.plot()\n        print()\n        init=False\n    print(f'---------fold{fold}---------')\n    print(f'train:{num_train} valid:{num_valid}')\n    print()\n    \n    if CFG.resume:\n        print(f'resume from epoch{CFG.resume}')\n        model.load_weights(f'{CFG.output_dir}/{CFG.comment}-fold{fold}-last.h5')\n        if train_ds is not None:\n            model.evaluate(train_ds.take(steps_per_epoch))\n        if valid_ds is not None:\n            model.evaluate(valid_ds)\n\n    logger = tf.keras.callbacks.CSVLogger(f'{CFG.output_dir}/{CFG.comment}-fold{fold}-logs.csv')\n    sv_loss = tf.keras.callbacks.ModelCheckpoint(f'{CFG.output_dir}/{CFG.comment}-fold{fold}-best.h5', monitor='val_loss', verbose=0, save_best_only=True,\n                save_weights_only=True, mode='min', save_freq='epoch')\n    snap = Snapshot(f'{CFG.output_dir}/{CFG.comment}-fold{fold}', CFG.snapshot_epochs)\n    swa = SWA(f'{CFG.output_dir}/{CFG.comment}-fold{fold}', CFG.swa_epochs, strategy=strategy, train_ds=train_ds, valid_ds=valid_ds, valid_steps=-(num_valid//-CFG.batch_size))\n    callbacks = []\n    if CFG.save_output:\n        callbacks.append(logger)\n        callbacks.append(snap)\n        callbacks.append(swa)\n        if fold != 'all':\n            callbacks.append(sv_loss)\n        \n    history = model.fit(\n        train_ds,\n        epochs=CFG.epoch-CFG.resume,\n        steps_per_epoch=steps_per_epoch,\n        callbacks=callbacks,\n        validation_data=valid_ds,\n        verbose=CFG.verbose,\n        validation_steps=-(num_valid//-CFG.batch_size)\n    )\n\n    if CFG.save_output:\n        try:\n            model.load_weights(f'{CFG.output_dir}/{CFG.comment}-fold{fold}-best.h5')\n        except:\n            pass\n    if fold != 'all':\n        cv = model.evaluate(valid_ds,verbose=CFG.verbose,steps=-(num_valid//-CFG.batch_size))\n    else:\n        cv = None\n\n    return model, cv, history\n\ndef train_folds(CFG, folds, strategy=STRATEGY, summary=True):\n    for fold in folds:\n        if fold != 'all':\n            all_files = TRAIN_FILENAMES\n            train_files = [x for x in all_files if f'fold{fold}' not in x]\n            valid_files = [x for x in all_files if f'fold{fold}' in x]\n        else:\n            train_files = TRAIN_FILENAMES\n            valid_files = None\n        \n        train_fold(CFG, fold, train_files, valid_files, strategy=strategy, summary=summary)\n    return","metadata":{"id":"BoGUEL6-oEWO","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# train_folds(CFG, [0])","metadata":{"id":"V8T5GhYPqAFR","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# CFG.seed = 42\n# CFG.comment = f'islr-fp16-192-8-seed{CFG.seed}'\n# train_folds(CFG, ['all'], summary=False)","metadata":{"id":"WTJ8eruttPxJ","outputId":"d5bdfb3b-a64d-4434-b176-f1d58c2adaa3","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# CFG.seed = 43\n# CFG.comment = f'islr-fp16-192-8-seed{CFG.seed}'\n# train_folds(CFG, ['all'], summary=False)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# CFG.seed = 44\n# CFG.comment = f'islr-fp16-192-8-seed{CFG.seed}'\n# train_folds(CFG, ['all'], summary=False)","metadata":{"id":"rHLmUWM4tP9S","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# CFG.seed = 45\n# CFG.comment = f'islr-fp16-192-8-seed{CFG.seed}'\n# train_folds(CFG, ['all'], summary=False)","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}