{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"},"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":"markdown","source":"**It takes around 6~7 hours to train a model with single fold.**\n\nSo, you should run this notebook four time (for each seed=42,43,44,45) to get all 4 seed weights of the model.\n\nOr you can just press 'Save and Run All' to get the result of the fold 0 model.\n\nThere are two issues with this notebook:\n\n- Training time issue: The weird thing is that it takes around 3 hours in Colab TPU (v2-8), which is supposed to be slower than Kaggle's TPU (v3-8). If you know a lot about tf+TPU frameworks, please let me know how to debug this issue.\n\n- Training unstability: I Changed some minor configurations from the final solution. I changed epoch 400 -> 300, clipvalue=1. -> None, label_smoothing=0 -> label_smoothing=0.1 for more stable reproducibility(with slightly lower accuracy). unstability mainly caused by high lambda value of the AWP(0.2) - it will cause nan loss somtimes(around once out of five times). you can switch it to 0.1 or lower the learning rate and still can get fairly high accuracy. Also, if you have any idea related to this issue, please let me know.","metadata":{}},{"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","execution":{"iopub.status.busy":"2023-05-03T18:59:07.331438Z","iopub.execute_input":"2023-05-03T18:59:07.332383Z","iopub.status.idle":"2023-05-03T18:59:23.706113Z","shell.execute_reply.started":"2023-05-03T18:59:07.332339Z","shell.execute_reply":"2023-05-03T18:59:23.704923Z"},"trusted":true},"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\n\nprint(f'Tensorflow Version: {tf.__version__}')\nprint(f'Python Version: {sys.version}')","metadata":{"id":"wD7tqFC_qAFQ","outputId":"426ff58f-01d7-451e-ea99-7b00405d8781","execution":{"iopub.status.busy":"2023-05-03T18:59:23.708109Z","iopub.execute_input":"2023-05-03T18:59:23.708406Z","iopub.status.idle":"2023-05-03T19:00:05.71278Z","shell.execute_reply.started":"2023-05-03T18:59:23.708378Z","shell.execute_reply":"2023-05-03T19:00:05.711747Z"},"trusted":true},"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","execution":{"iopub.status.busy":"2023-05-03T19:00:05.714174Z","iopub.execute_input":"2023-05-03T19:00:05.714749Z","iopub.status.idle":"2023-05-03T19:00:16.235478Z","shell.execute_reply.started":"2023-05-03T19:00:05.714719Z","shell.execute_reply":"2023-05-03T19:00:16.234303Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"TRAIN_FILENAMES = glob.glob('/kaggle/input/islr-5fold/*.tfrecords')\nprint(len(TRAIN_FILENAMES))","metadata":{"id":"QVDc3vkwqAFQ","outputId":"b84108e5-9e50-4a46-eca8-6a60ecdde99c","execution":{"iopub.status.busy":"2023-05-03T19:01:19.125658Z","iopub.execute_input":"2023-05-03T19:01:19.126773Z","iopub.status.idle":"2023-05-03T19:01:19.183316Z","shell.execute_reply.started":"2023-05-03T19:01:19.126732Z","shell.execute_reply":"2023-05-03T19:01:19.182241Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Train DataFrame\ntrain_df = pd.read_csv('/kaggle/input/asl-signs/train.csv')\ndisplay(train_df.head())\ndisplay(train_df.info())","metadata":{"id":"DtrBI-jwqAFQ","outputId":"3ebb3c2b-8fbe-4552-8a16-d50484823b57","execution":{"iopub.status.busy":"2023-05-03T19:01:20.089446Z","iopub.execute_input":"2023-05-03T19:01:20.089916Z","iopub.status.idle":"2023-05-03T19:01:20.309079Z","shell.execute_reply.started":"2023-05-03T19:01:20.089882Z","shell.execute_reply":"2023-05-03T19:01:20.308065Z"},"trusted":true},"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)\nprint(count_data_items(TRAIN_FILENAMES), len(train_df))\nassert count_data_items(TRAIN_FILENAMES) == len(train_df)","metadata":{"id":"iAc9FKztqAFQ","outputId":"ade0413c-5a17-4a4f-f5a6-b0048ab76ffa","execution":{"iopub.status.busy":"2023-05-03T19:01:20.310983Z","iopub.execute_input":"2023-05-03T19:01:20.311328Z","iopub.status.idle":"2023-05-03T19:01:20.319186Z","shell.execute_reply.started":"2023-05-03T19:01:20.311298Z","shell.execute_reply":"2023-05-03T19:01:20.318256Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ROWS_PER_FRAME = 543\nMAX_LEN = 384\nCROP_LEN = MAX_LEN\nNUM_CLASSES  = 250\nPAD = -100.\nNOSE=[\n    1,2,98,327\n]\nLNOSE = [98]\nRNOSE = [327]\nLIP = [ 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]\nLLIP = [84,181,91,146,61,185,40,39,37,87,178,88,95,78,191,80,81,82]\nRLIP = [314,405,321,375,291,409,270,269,267,317,402,318,324,308,415,310,311,312]\n\nPOSE = [500, 502, 504, 501, 503, 505, 512, 513]\nLPOSE = [513,505,503,501]\nRPOSE = [512,504,502,500]\n\nREYE = [\n    33, 7, 163, 144, 145, 153, 154, 155, 133,\n    246, 161, 160, 159, 158, 157, 173,\n]\nLEYE = [\n    263, 249, 390, 373, 374, 380, 381, 382, 362,\n    466, 388, 387, 386, 385, 384, 398,\n]\n\nLHAND = np.arange(468, 489).tolist()\nRHAND = np.arange(522, 543).tolist()\n\nPOINT_LANDMARKS = LIP + LHAND + RHAND + NOSE + REYE + LEYE #+POSE\n\nNUM_NODES = len(POINT_LANDMARKS)\nCHANNELS = 6*NUM_NODES\n\nprint(NUM_NODES)\nprint(CHANNELS)\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","execution":{"iopub.status.busy":"2023-05-03T19:01:20.324347Z","iopub.execute_input":"2023-05-03T19:01:20.324636Z","iopub.status.idle":"2023-05-03T19:01:20.361019Z","shell.execute_reply.started":"2023-05-03T19:01:20.324602Z","shell.execute_reply":"2023-05-03T19:01:20.359852Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def 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\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\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\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\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\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\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\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\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\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\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        # 使用 pad_batch 并明确指定填充值和形状\n        ds = ds.padded_batch(\n            batch_size, \n            padding_values=(PAD, 0.0),  # 对样本和标签分别设置填充值\n            padded_shapes=(\n                [max_len, CHANNELS],  # 样本形状\n                [NUM_CLASSES]  # 标签形状\n            ),\n            drop_remainder=drop_remainder\n        )\n    ds = ds.prefetch(tf.data.AUTOTUNE)\n        \n    return ds\n","metadata":{"id":"r17ZnZaGqAFQ","execution":{"iopub.status.busy":"2023-05-03T19:01:20.483306Z","iopub.execute_input":"2023-05-03T19:01:20.483604Z","iopub.status.idle":"2023-05-03T19:01:25.210581Z","shell.execute_reply.started":"2023-05-03T19:01:20.483566Z","shell.execute_reply":"2023-05-03T19:01:25.209138Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def efficient_cutmix(samples, labels, alpha=0.7, cutmix_ratio=0.2):\n    \"\"\"\n    高效的CutMix实现，使用向量化操作代替循环\n    适用于手语识别的时序数据\n    \n    参数:\n        samples: 输入特征 [batch_size, time_steps, channels]\n        labels: 标签 [batch_size, num_classes]\n        alpha: 混合权重\n        cutmix_ratio: 要处理的批次中样本的比例\n    \"\"\"\n    batch_size = tf.shape(samples)[0]\n    \n    # 仅在批次大小足够时执行CutMix\n    def do_mix():\n        # 创建随机排列的索引\n        perm = tf.random.shuffle(tf.range(batch_size))\n        \n        # 创建CutMix掩码 - 决定哪些样本会被混合\n        cutmix_samples = tf.cast(tf.random.uniform([batch_size]) < cutmix_ratio, tf.float32)\n        cutmix_samples = tf.reshape(cutmix_samples, [-1, 1, 1])  # 广播到时间和特征维度\n        \n        # 创建手部和嘴唇区域的掩码\n        # 预先计算这些掩码可以避免每个样本都执行复杂的索引操作\n        channels = tf.shape(samples)[2]\n        feature_mask = tf.zeros([1, 1, channels], dtype=tf.float32)\n        \n        # 构建简化的特征掩码 - 我们将样本分为三部分：\n        # 1. 需要重点混合的区域 (掩码值 = 1.0)\n        # 2. 轻度混合的区域 (掩码值 = 0.5)\n        # 3. 保持原样的区域 (掩码值 = 0.0)\n        \n        # 随机生成一个三区域掩码，而不是精确针对手部/嘴唇\n        # 这样可以大大提高效率，同时仍然提供多样化的混合\n        random_mask = tf.random.uniform([1, 1, channels], 0, 1)\n        \n        # 创建三区域掩码\n        high_mix_mask = tf.cast(random_mask < 0.3, tf.float32)  # 30% 通道重点混合\n        medium_mix_mask = tf.cast((random_mask >= 0.3) & (random_mask < 0.6), tf.float32)  # 30% 通道轻度混合\n        \n        # 混合样本和置换样本\n        mixed_samples = (\n            (1 - cutmix_samples * high_mix_mask * 0.8) * samples +  # 原始样本，高混合区域减少80%\n            (cutmix_samples * high_mix_mask * 0.8) * tf.gather(samples, perm) +  # 置换样本，高混合区域占80%\n            \n            (1 - cutmix_samples * medium_mix_mask * 0.4) * samples +  # 原始样本，中混合区域减少40%\n            (cutmix_samples * medium_mix_mask * 0.4) * tf.gather(samples, perm)  # 置换样本，中混合区域占40%\n        )\n        \n        # 混合标签 - 使用简单的线性插值\n        cutmix_labels = tf.reshape(cutmix_samples, [-1, 1])  # 广播到标签维度\n        mixed_labels = (\n            (1 - cutmix_labels * alpha) * labels +\n            (cutmix_labels * alpha) * tf.gather(labels, perm)\n        )\n        \n        return mixed_samples, mixed_labels\n    \n    def no_mix():\n        return samples, labels\n    \n    # 只有当批次大小 > 1 时才执行CutMix\n    return tf.cond(batch_size > 1, do_mix, no_mix)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def augment_dataset_with_cutmix(dataset, cutmix_prob=0.5, cutmix_ratio=0.2, alpha=0.7):\n    \"\"\"\n    对数据集应用CutMix增强\n    \"\"\"\n    def apply_cutmix(samples, labels):\n        # 随机决定是否应用CutMix\n        should_apply = tf.random.uniform(shape=[], dtype=tf.float32) < cutmix_prob\n        \n        def do_cutmix():\n            return efficient_cutmix(samples, labels, alpha=alpha, cutmix_ratio=cutmix_ratio)\n        \n        def no_cutmix():\n            return samples, labels\n        \n        return tf.cond(should_apply, do_cutmix, no_cutmix)\n    \n    # 使用map应用CutMix函数到每个批次\n    return dataset.map(apply_cutmix, num_parallel_calls=tf.data.AUTOTUNE)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from IPython.display import HTML\nimport matplotlib.animation as animation\nfrom matplotlib.animation import FuncAnimation\n\n#调试信息\n#调用数据增强\nds = get_tfrec_dataset(TRAIN_FILENAMES, augment=True, batch_size=1024)\n# 打印数据集信息\nfor samples, labels in ds.take(1):\n    print(\"Dataset samples shape:\", samples.shape)\n    print(\"Dataset labels shape:\", labels.shape)\n\nds_cutmix = augment_dataset_with_cutmix(ds, cutmix_prob=0.5, cutmix_ratio=0.2)\n\n    \nedges = [(0,1),(1,2),(2,3),(3,4),(0,5),(0,17),(5,6),(6,7),(7,8),(5,9),(9,10),(10,11),(11,12),\n         (9,13),(13,14),(14,15),(15,16),(13,17),(17,18),(18,19),(19,20)]\n\n\ndef plot_frame(frame, edges=[], idxs=[]):\n        \n    frame[np.isnan(frame)] = 0\n    x = list(frame[...,0])\n    y = list(frame[...,1])\n    if len(idxs) == 0:\n        idxs = list(range(len(x)))\n    ax.clear()\n    ax.scatter(x, y, color='dodgerblue')\n    for i in range(len(x)):\n        ax.text(x[i], y[i], idxs[i])\n        \n    for edge in edges:\n        ax.plot([x[edge[0]], x[edge[1]]], [y[edge[0]], y[edge[1]]], color='salmon')\n    ax.set_xticks([])\n    ax.set_yticks([])\n    ax.set_xticklabels([])\n    ax.set_yticklabels([])\n\ndef animate_frames(frames, edges=[], idxs=[]):\n    anim = FuncAnimation(fig, lambda frame: plot_frame(frame, edges, idxs), frames=frames, interval=100)\n    return HTML(anim.to_jshtml())","metadata":{"id":"gbVSTH9xqAFQ","outputId":"c212e69c-230c-4ded-b204-183dd28eebcb","execution":{"iopub.status.busy":"2023-05-03T19:01:25.213012Z","iopub.execute_input":"2023-05-03T19:01:25.213364Z","iopub.status.idle":"2023-05-03T19:01:25.681328Z","shell.execute_reply.started":"2023-05-03T19:01:25.213333Z","shell.execute_reply":"2023-05-03T19:01:25.6803Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class 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\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\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\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","execution":{"iopub.status.busy":"2023-05-03T19:02:52.737995Z","iopub.execute_input":"2023-05-03T19:02:52.738436Z","iopub.status.idle":"2023-05-03T19:02:52.766619Z","shell.execute_reply.started":"2023-05-03T19:02:52.738404Z","shell.execute_reply":"2023-05-03T19:02:52.765627Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class 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\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","execution":{"iopub.status.busy":"2023-05-03T19:02:52.767813Z","iopub.execute_input":"2023-05-03T19:02:52.768346Z","iopub.status.idle":"2023-05-03T19:02:52.789973Z","shell.execute_reply.started":"2023-05-03T19:02:52.768308Z","shell.execute_reply":"2023-05-03T19:02:52.788959Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def 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)\n","metadata":{"id":"KIooIcnSqAFR","outputId":"4d6a1f52-1630-45d6-a33c-441aab02d1b2","execution":{"iopub.status.busy":"2023-05-03T19:02:52.792887Z","iopub.execute_input":"2023-05-03T19:02:52.79325Z","iopub.status.idle":"2023-05-03T19:02:57.052428Z","shell.execute_reply.started":"2023-05-03T19:02:52.793225Z","shell.execute_reply":"2023-05-03T19:02:57.051468Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 首先定义一个自定义回调函数，用于根据当前epoch更新数据集\nclass DatasetUpdater(tf.keras.callbacks.Callback):\n    def __init__(self, train_files, CFG, cutmix_epoch_start=50):\n        super().__init__()\n        self.train_files = train_files\n        self.CFG = CFG\n        self.cutmix_epoch_start = cutmix_epoch_start\n        # 保存初始数据集以供第一个epoch使用\n        self.initial_ds = get_tfrec_dataset(\n            train_files, \n            batch_size=CFG.batch_size, \n            max_len=CFG.max_len, \n            drop_remainder=True if 'all' != 'all' else False, \n            augment=True, \n            repeat=True, \n            shuffle=32768\n        )\n        \n    def on_epoch_begin(self, epoch, logs=None):\n        # 注意：epoch从0开始计数\n        if epoch + 1 >= self.cutmix_epoch_start:  # +1是为了按照人类习惯计数epoch\n            # 创建带有CutMix的数据集\n            ds = get_tfrec_dataset(\n                self.train_files, \n                batch_size=self.CFG.batch_size, \n                max_len=self.CFG.max_len, \n                drop_remainder=True if 'all' != 'all' else False, \n                augment=True, \n                repeat=True, \n                shuffle=32768\n            )\n            #ds = augment_dataset_with_cutmix(ds, cutmix_prob=0.2, cutmix_ratio=0.3, alpha=0.7)\n            ds = augment_dataset_with_cutmix(ds, cutmix_prob=0.2, cutmix_ratio=0.15, alpha=0.3)\n            print(f\"Epoch {epoch+1}: Using CutMix augmentation\")\n        else:\n            # 使用没有CutMix的数据集\n            ds = get_tfrec_dataset(\n                self.train_files, \n                batch_size=self.CFG.batch_size, \n                max_len=self.CFG.max_len, \n                drop_remainder=True if 'all' != 'all' else False, \n                augment=True, \n                repeat=True, \n                shuffle=32768\n            )\n            print(f\"Epoch {epoch+1}: Not using CutMix augmentation\")\n            \n        # 更新模型的训练数据集\n        self.model.train_dataset = ds","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_fold(CFG, fold, train_files, valid_files=None, strategy=STRATEGY, summary=True, cutmix_epoch_start=50):\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    # 创建初始训练数据集（不带CutMix）\n    if fold != 'all':\n        train_ds = get_tfrec_dataset(\n            train_files, \n            batch_size=CFG.batch_size, \n            max_len=CFG.max_len, \n            drop_remainder=True, \n            augment=True, \n            repeat=True, \n            shuffle=32768\n        )\n        \n        valid_ds = get_tfrec_dataset(\n            valid_files, \n            batch_size=CFG.batch_size, \n            max_len=CFG.max_len, \n            drop_remainder=False, \n            repeat=False, \n            shuffle=False\n        )\n    else:\n        train_ds = get_tfrec_dataset(\n            train_files, \n            batch_size=CFG.batch_size, \n            max_len=CFG.max_len, \n            drop_remainder=False, \n            augment=True, \n            repeat=True, \n            shuffle=32768\n        )\n        valid_ds = None\n        valid_files = []\n    \n    # 创建数据集更新器回调\n    dataset_updater = DatasetUpdater(train_files, CFG, cutmix_epoch_start)\n    \n    num_train = count_data_items(train_files)\n    num_valid = count_data_items(valid_files) if valid_files else 0\n    steps_per_epoch = num_train//CFG.batch_size\n    \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)],\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    \n    # 添加回调\n    callbacks = [dataset_updater]  # 添加数据集更新器回调\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) if num_valid > 0 else None\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\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","execution":{"iopub.status.busy":"2023-05-03T19:02:57.060277Z","iopub.execute_input":"2023-05-03T19:02:57.060615Z","iopub.status.idle":"2023-05-03T19:02:57.095753Z","shell.execute_reply.started":"2023-05-03T19:02:57.060575Z","shell.execute_reply":"2023-05-03T19:02:57.094879Z"},"trusted":true},"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 #400\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","execution":{"iopub.status.busy":"2023-05-03T19:02:57.096792Z","iopub.execute_input":"2023-05-03T19:02:57.097097Z","iopub.status.idle":"2023-05-03T19:02:57.122726Z","shell.execute_reply.started":"2023-05-03T19:02:57.09707Z","shell.execute_reply":"2023-05-03T19:02:57.121838Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 定义训练和验证文件\nall_files = TRAIN_FILENAMES\nfold = 0  # 选择要训练的折\ntrain_files = [x for x in all_files if f'fold{fold}' not in x]\nvalid_files = [x for x in all_files if f'fold{fold}' in x]\n\n# 然后调用训练函数\ntrain_fold(CFG, fold, train_files, valid_files, cutmix_epoch_start=0)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"id":"V8T5GhYPqAFR","execution":{"iopub.status.busy":"2023-05-03T19:02:57.123783Z","iopub.execute_input":"2023-05-03T19:02:57.124086Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#uncomment each cell and comment out all the other cells including train_folds to get the result","metadata":{},"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","execution":{"iopub.status.busy":"2023-05-03T19:00:17.178358Z","iopub.status.idle":"2023-05-03T19:00:17.17871Z","shell.execute_reply.started":"2023-05-03T19:00:17.178515Z","shell.execute_reply":"2023-05-03T19:00:17.178531Z"},"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":{},"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","execution":{"iopub.status.busy":"2023-05-03T19:00:17.180301Z","iopub.status.idle":"2023-05-03T19:00:17.180698Z","shell.execute_reply.started":"2023-05-03T19:00:17.180484Z","shell.execute_reply":"2023-05-03T19:00:17.180501Z"},"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":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"id":"at4Io1dNtSLh"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"id":"pRnjnF6Cj9nY"},"outputs":[],"execution_count":null}]}