两道新题
This commit is contained in:
@@ -4,7 +4,11 @@
|
|||||||
"cell_type": "code",
|
"cell_type": "code",
|
||||||
"id": "initial_id",
|
"id": "initial_id",
|
||||||
"metadata": {
|
"metadata": {
|
||||||
"collapsed": true
|
"collapsed": true,
|
||||||
|
"ExecuteTime": {
|
||||||
|
"end_time": "2026-07-05T15:05:59.049967188Z",
|
||||||
|
"start_time": "2026-07-05T15:05:56.833387637Z"
|
||||||
|
}
|
||||||
},
|
},
|
||||||
"source": [
|
"source": [
|
||||||
"import os\n",
|
"import os\n",
|
||||||
@@ -17,10 +21,15 @@
|
|||||||
"from tqdm import tqdm"
|
"from tqdm import tqdm"
|
||||||
],
|
],
|
||||||
"outputs": [],
|
"outputs": [],
|
||||||
"execution_count": null
|
"execution_count": 2
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
"metadata": {},
|
"metadata": {
|
||||||
|
"ExecuteTime": {
|
||||||
|
"end_time": "2026-07-05T15:05:59.150289671Z",
|
||||||
|
"start_time": "2026-07-05T15:05:59.051997256Z"
|
||||||
|
}
|
||||||
|
},
|
||||||
"cell_type": "code",
|
"cell_type": "code",
|
||||||
"source": [
|
"source": [
|
||||||
"train_df = pd.read_csv('./train.csv')\n",
|
"train_df = pd.read_csv('./train.csv')\n",
|
||||||
@@ -30,16 +39,136 @@
|
|||||||
"print(f\"Train: {len(train_df)}, Val: {len(val_df)}, Test: {len(test_df)}\")"
|
"print(f\"Train: {len(train_df)}, Val: {len(val_df)}, Test: {len(test_df)}\")"
|
||||||
],
|
],
|
||||||
"id": "f32065752f1597fb",
|
"id": "f32065752f1597fb",
|
||||||
"outputs": [],
|
"outputs": [
|
||||||
"execution_count": null
|
{
|
||||||
|
"name": "stdout",
|
||||||
|
"output_type": "stream",
|
||||||
|
"text": [
|
||||||
|
"Train: 10908, Val: 1212, Test: 5195\n"
|
||||||
|
]
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"execution_count": 3
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
"metadata": {},
|
"metadata": {
|
||||||
|
"ExecuteTime": {
|
||||||
|
"end_time": "2026-07-05T15:05:59.184794043Z",
|
||||||
|
"start_time": "2026-07-05T15:05:59.156968452Z"
|
||||||
|
}
|
||||||
|
},
|
||||||
"cell_type": "code",
|
"cell_type": "code",
|
||||||
"source": "train_df.head()",
|
"source": "train_df.head()",
|
||||||
"id": "460c6fffc81814df",
|
"id": "460c6fffc81814df",
|
||||||
"outputs": [],
|
"outputs": [
|
||||||
"execution_count": null
|
{
|
||||||
|
"data": {
|
||||||
|
"text/plain": [
|
||||||
|
" id premise \\\n",
|
||||||
|
"9989 6417481c8d यद्यपि हम आज के समय में अल कायदा के साथ केएसएए... \n",
|
||||||
|
"3880 5c5ca34cf6 'Upload him into his body? What body?' \n",
|
||||||
|
"8559 1d3c28ecff yeah and then about every five years you have ... \n",
|
||||||
|
"6316 22e2a4903d θέλω να πω ότι υπήρχε είχα, είχα το ρολόι μου ... \n",
|
||||||
|
"762 307016c21f Yet, in the mouths of the white townsfolk of S... \n",
|
||||||
|
"\n",
|
||||||
|
" hypothesis lang_abv language \\\n",
|
||||||
|
"9989 हर व्यक्ति केएसएम को हमेशा अल कायदा के बराबर म... hi Hindi \n",
|
||||||
|
"3880 I don't think he has a body at all. en English \n",
|
||||||
|
"8559 You have to dig them up every five years, thro... en English \n",
|
||||||
|
"6316 Ήταν ενοχλητικό όταν κάλυψε τα παπούτσια μου. el Greek \n",
|
||||||
|
"762 White townsfolk in Salisbury, N.C. are easily ... en English \n",
|
||||||
|
"\n",
|
||||||
|
" label \n",
|
||||||
|
"9989 2 \n",
|
||||||
|
"3880 1 \n",
|
||||||
|
"8559 1 \n",
|
||||||
|
"6316 1 \n",
|
||||||
|
"762 1 "
|
||||||
|
],
|
||||||
|
"text/html": [
|
||||||
|
"<div>\n",
|
||||||
|
"<style scoped>\n",
|
||||||
|
" .dataframe tbody tr th:only-of-type {\n",
|
||||||
|
" vertical-align: middle;\n",
|
||||||
|
" }\n",
|
||||||
|
"\n",
|
||||||
|
" .dataframe tbody tr th {\n",
|
||||||
|
" vertical-align: top;\n",
|
||||||
|
" }\n",
|
||||||
|
"\n",
|
||||||
|
" .dataframe thead th {\n",
|
||||||
|
" text-align: right;\n",
|
||||||
|
" }\n",
|
||||||
|
"</style>\n",
|
||||||
|
"<table border=\"1\" class=\"dataframe\">\n",
|
||||||
|
" <thead>\n",
|
||||||
|
" <tr style=\"text-align: right;\">\n",
|
||||||
|
" <th></th>\n",
|
||||||
|
" <th>id</th>\n",
|
||||||
|
" <th>premise</th>\n",
|
||||||
|
" <th>hypothesis</th>\n",
|
||||||
|
" <th>lang_abv</th>\n",
|
||||||
|
" <th>language</th>\n",
|
||||||
|
" <th>label</th>\n",
|
||||||
|
" </tr>\n",
|
||||||
|
" </thead>\n",
|
||||||
|
" <tbody>\n",
|
||||||
|
" <tr>\n",
|
||||||
|
" <th>9989</th>\n",
|
||||||
|
" <td>6417481c8d</td>\n",
|
||||||
|
" <td>यद्यपि हम आज के समय में अल कायदा के साथ केएसएए...</td>\n",
|
||||||
|
" <td>हर व्यक्ति केएसएम को हमेशा अल कायदा के बराबर म...</td>\n",
|
||||||
|
" <td>hi</td>\n",
|
||||||
|
" <td>Hindi</td>\n",
|
||||||
|
" <td>2</td>\n",
|
||||||
|
" </tr>\n",
|
||||||
|
" <tr>\n",
|
||||||
|
" <th>3880</th>\n",
|
||||||
|
" <td>5c5ca34cf6</td>\n",
|
||||||
|
" <td>'Upload him into his body? What body?'</td>\n",
|
||||||
|
" <td>I don't think he has a body at all.</td>\n",
|
||||||
|
" <td>en</td>\n",
|
||||||
|
" <td>English</td>\n",
|
||||||
|
" <td>1</td>\n",
|
||||||
|
" </tr>\n",
|
||||||
|
" <tr>\n",
|
||||||
|
" <th>8559</th>\n",
|
||||||
|
" <td>1d3c28ecff</td>\n",
|
||||||
|
" <td>yeah and then about every five years you have ...</td>\n",
|
||||||
|
" <td>You have to dig them up every five years, thro...</td>\n",
|
||||||
|
" <td>en</td>\n",
|
||||||
|
" <td>English</td>\n",
|
||||||
|
" <td>1</td>\n",
|
||||||
|
" </tr>\n",
|
||||||
|
" <tr>\n",
|
||||||
|
" <th>6316</th>\n",
|
||||||
|
" <td>22e2a4903d</td>\n",
|
||||||
|
" <td>θέλω να πω ότι υπήρχε είχα, είχα το ρολόι μου ...</td>\n",
|
||||||
|
" <td>Ήταν ενοχλητικό όταν κάλυψε τα παπούτσια μου.</td>\n",
|
||||||
|
" <td>el</td>\n",
|
||||||
|
" <td>Greek</td>\n",
|
||||||
|
" <td>1</td>\n",
|
||||||
|
" </tr>\n",
|
||||||
|
" <tr>\n",
|
||||||
|
" <th>762</th>\n",
|
||||||
|
" <td>307016c21f</td>\n",
|
||||||
|
" <td>Yet, in the mouths of the white townsfolk of S...</td>\n",
|
||||||
|
" <td>White townsfolk in Salisbury, N.C. are easily ...</td>\n",
|
||||||
|
" <td>en</td>\n",
|
||||||
|
" <td>English</td>\n",
|
||||||
|
" <td>1</td>\n",
|
||||||
|
" </tr>\n",
|
||||||
|
" </tbody>\n",
|
||||||
|
"</table>\n",
|
||||||
|
"</div>"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
"execution_count": 4,
|
||||||
|
"metadata": {},
|
||||||
|
"output_type": "execute_result"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"execution_count": 4
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
"metadata": {
|
"metadata": {
|
||||||
|
|||||||
File diff suppressed because one or more lines are too long
@@ -0,0 +1,221 @@
|
|||||||
|
# %%
|
||||||
|
import tensorflow as tf
|
||||||
|
import torch
|
||||||
|
import torchvision
|
||||||
|
from torch.utils.data import IterableDataset, DataLoader
|
||||||
|
from matplotlib import pyplot as plt
|
||||||
|
import numpy as np
|
||||||
|
import torch.nn as nn
|
||||||
|
from PIL import Image
|
||||||
|
def parse_tfrecord(example_proto):
|
||||||
|
feature_description = {
|
||||||
|
'image': tf.io.FixedLenFeature([], tf.string),
|
||||||
|
'class': tf.io.FixedLenFeature([], tf.int64),
|
||||||
|
'id' : tf.io.FixedLenFeature([], tf.string),
|
||||||
|
}
|
||||||
|
parsed = tf.io.parse_single_example(example_proto, feature_description)
|
||||||
|
image = tf.image.decode_jpeg(parsed['image'], channels=3)
|
||||||
|
image = tf.image.resize(image, [224, 224])
|
||||||
|
#image = tf.image.convert_image_dtype(image, tf.float32)
|
||||||
|
label = parsed['class']
|
||||||
|
idd = parsed['id']
|
||||||
|
return image, label,idd
|
||||||
|
|
||||||
|
def load_tfrecord_dataset(pattern):
|
||||||
|
files = tf.io.gfile.glob(pattern)
|
||||||
|
if not files:
|
||||||
|
raise ValueError(f"No files found for pattern {pattern}")
|
||||||
|
dataset = tf.data.TFRecordDataset(files)
|
||||||
|
dataset = dataset.map(parse_tfrecord)
|
||||||
|
# 可选:打乱、批处理等,但此处我们只返回样本级别的数据集
|
||||||
|
return dataset
|
||||||
|
|
||||||
|
class TFRecordToPyTorch(IterableDataset):
|
||||||
|
def __init__(self, tfrecord_pattern,transform=None):
|
||||||
|
self.tfrecord_pattern = tfrecord_pattern
|
||||||
|
self.transform=transform
|
||||||
|
|
||||||
|
def __iter__(self):
|
||||||
|
# 每次迭代创建新的数据集,保证可重复使用
|
||||||
|
dataset = load_tfrecord_dataset(self.tfrecord_pattern)
|
||||||
|
# 使用 as_numpy_iterator() 获取 NumPy 数组,便于转换为 PyTorch 张量
|
||||||
|
for image_np, label_np,idd in dataset.as_numpy_iterator():
|
||||||
|
# image_np shape: (224,224,3), dtype float32, label_np scalar int64
|
||||||
|
# 转为 PyTorch 张量,并调整为 CxHxW
|
||||||
|
image_pil = Image.fromarray((image_np).astype('uint8'))
|
||||||
|
if self.transform:
|
||||||
|
image_tensor = self.transform(image_pil)
|
||||||
|
else:
|
||||||
|
# 如果不需要 transform,至少转为 tensor
|
||||||
|
image_tensor = torch.from_numpy(image_np).permute(2,0,1)
|
||||||
|
#image_torch = torch.from_numpy(image_np).permute(2, 0, 1) # (3,224,224)
|
||||||
|
label_torch = torch.tensor(label_np, dtype=torch.long)
|
||||||
|
id_torch = idd
|
||||||
|
yield image_tensor, label_torch,id_torch
|
||||||
|
|
||||||
|
# 使用
|
||||||
|
transform = torchvision.transforms.Compose([
|
||||||
|
torchvision.transforms.ToTensor(),
|
||||||
|
torchvision.transforms.Normalize(mean=[0.485, 0.456, 0.406],
|
||||||
|
std=[0.229, 0.224, 0.225])
|
||||||
|
])
|
||||||
|
tfrecord_path = '/kaggle/input/competitions/tpu-getting-started/tfrecords-jpeg-224x224/train/*'
|
||||||
|
dataset = TFRecordToPyTorch(tfrecord_path,transform)
|
||||||
|
tfrecord_path = '/kaggle/input/competitions/tpu-getting-started/tfrecords-jpeg-224x224/val/*'
|
||||||
|
dataset2 = TFRecordToPyTorch(tfrecord_path,transform)
|
||||||
|
# 可以配合 DataLoader 使用
|
||||||
|
train_dataloader = DataLoader(dataset, batch_size=32, num_workers=0) # num_workers 设为0,因为 TF 数据集内部已并行
|
||||||
|
val_dataloader = DataLoader(dataset2, batch_size=32, num_workers=0)
|
||||||
|
for batch in train_dataloader:
|
||||||
|
plt.imshow(batch[0][1].permute(1,2,0).numpy())
|
||||||
|
break
|
||||||
|
plt.axis('off')
|
||||||
|
plt.show()
|
||||||
|
# %%
|
||||||
|
pretrained_net = torchvision.models.resnet50(pretrained=True)
|
||||||
|
# %%
|
||||||
|
pretrained_net.fc=nn.Linear(pretrained_net.fc.in_features,104)
|
||||||
|
nn.init.xavier_uniform_(pretrained_net.fc.weight)
|
||||||
|
# %%
|
||||||
|
for parm in pretrained_net.conv1.parameters():
|
||||||
|
parm.requires_grad=False
|
||||||
|
for parm in pretrained_net.bn1.parameters():
|
||||||
|
parm.requires_grad=False
|
||||||
|
for parm in pretrained_net.layer1.parameters():
|
||||||
|
parm.requires_grad=False
|
||||||
|
for parm in pretrained_net.layer2.parameters():
|
||||||
|
parm.requires_grad=False
|
||||||
|
for parm in pretrained_net.layer3.parameters():
|
||||||
|
parm.requires_grad=False
|
||||||
|
# %%
|
||||||
|
def print_trainable_info(model):
|
||||||
|
frozen = sum(p.numel() for p in model.parameters() if not p.requires_grad)
|
||||||
|
trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)
|
||||||
|
total = frozen + trainable
|
||||||
|
print(f" 冻结参数: {frozen:,} 可训练参数: {trainable:,} ({100.*trainable/total:.1f}%)")
|
||||||
|
# %%
|
||||||
|
print_trainable_info(pretrained_net)
|
||||||
|
# %%
|
||||||
|
@torch.no_grad()
|
||||||
|
def validate(model,loader):
|
||||||
|
model.eval()
|
||||||
|
acc=0
|
||||||
|
total=0
|
||||||
|
for batch in loader:
|
||||||
|
X = batch[0]
|
||||||
|
labels = batch[1]
|
||||||
|
X = X.to(device)
|
||||||
|
labels = labels.to(device)
|
||||||
|
pred=torch.argmax(model(X),dim=1)
|
||||||
|
acc+=pred.eq(labels).sum()
|
||||||
|
total+=labels.size(0)
|
||||||
|
print(f"acc:{acc/total}")
|
||||||
|
# %%
|
||||||
|
from tqdm import tqdm
|
||||||
|
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
|
||||||
|
pretrained_net=pretrained_net.to(device)
|
||||||
|
loss_func = nn.CrossEntropyLoss()
|
||||||
|
optimizer = torch.optim.AdamW(pretrained_net.parameters(), lr=2e-4)
|
||||||
|
scheduler=torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=10)
|
||||||
|
epochs = 20
|
||||||
|
for epoch in range(epochs):
|
||||||
|
pretrained_net.train()
|
||||||
|
training_loss = 0
|
||||||
|
|
||||||
|
# 使用 tqdm 包装 dataloader,并设置描述信息
|
||||||
|
progress_bar = tqdm(train_dataloader, desc=f"Epoch {epoch+1}/{epochs}")
|
||||||
|
lens=0
|
||||||
|
for batch in progress_bar:
|
||||||
|
optimizer.zero_grad()
|
||||||
|
X = batch[0].to(device)
|
||||||
|
labels = batch[1].to(device)
|
||||||
|
|
||||||
|
outputs = pretrained_net(X)
|
||||||
|
loss = loss_func(outputs, labels)
|
||||||
|
loss.backward()
|
||||||
|
optimizer.step()
|
||||||
|
|
||||||
|
training_loss += loss.item()
|
||||||
|
lens+=labels.size(0)
|
||||||
|
# 更新进度条显示当前 batch 的损失
|
||||||
|
progress_bar.set_postfix({
|
||||||
|
'loss': f'{loss.item():.4f}',
|
||||||
|
'avg_loss': f'{training_loss / (progress_bar.n+1):.4f}' # progress_bar.n 是已处理 batch 数
|
||||||
|
})
|
||||||
|
|
||||||
|
scheduler.step()
|
||||||
|
|
||||||
|
# 计算平均训练损失(注意:len(train_dataloader) 才是 batch 总数)
|
||||||
|
avg_train_loss = training_loss / lens
|
||||||
|
print(f"Epoch {epoch+1} train_loss: {avg_train_loss:.4f}")
|
||||||
|
|
||||||
|
# 验证(你也可以为验证添加进度条,见下方建议)
|
||||||
|
validate(pretrained_net, val_dataloader)
|
||||||
|
|
||||||
|
# %%
|
||||||
|
import pandas as pd
|
||||||
|
def parse_tfrecord_test(example_proto):
|
||||||
|
feature_description = {
|
||||||
|
'image': tf.io.FixedLenFeature([], tf.string),
|
||||||
|
'id' : tf.io.FixedLenFeature([], tf.string)
|
||||||
|
}
|
||||||
|
parsed = tf.io.parse_single_example(example_proto, feature_description)
|
||||||
|
image = tf.image.decode_jpeg(parsed['image'], channels=3)
|
||||||
|
image = tf.image.resize(image, [224, 224])
|
||||||
|
#image = tf.image.convert_image_dtype(image, tf.float32)
|
||||||
|
idd = parsed['id']
|
||||||
|
return image,idd
|
||||||
|
def load_tfrecord_dataset_test(pattern):
|
||||||
|
files = tf.io.gfile.glob(pattern)
|
||||||
|
if not files:
|
||||||
|
raise ValueError(f"No files found for pattern {pattern}")
|
||||||
|
dataset = tf.data.TFRecordDataset(files)
|
||||||
|
dataset = dataset.map(parse_tfrecord_test)
|
||||||
|
# 可选:打乱、批处理等,但此处我们只返回样本级别的数据集
|
||||||
|
return dataset
|
||||||
|
class TFRecordToPyTorchTest(IterableDataset):
|
||||||
|
def __init__(self, tfrecord_pattern,transform=None):
|
||||||
|
self.tfrecord_pattern = tfrecord_pattern
|
||||||
|
self.transform=transform
|
||||||
|
|
||||||
|
def __iter__(self):
|
||||||
|
# 每次迭代创建新的数据集,保证可重复使用
|
||||||
|
dataset = load_tfrecord_dataset_test(self.tfrecord_pattern)
|
||||||
|
# 使用 as_numpy_iterator() 获取 NumPy 数组,便于转换为 PyTorch 张量
|
||||||
|
for image_np,idd in dataset.as_numpy_iterator():
|
||||||
|
# image_np shape: (224,224,3), dtype float32, label_np scalar int64
|
||||||
|
# 转为 PyTorch 张量,并调整为 CxHxW
|
||||||
|
image_pil = Image.fromarray((image_np).astype('uint8'))
|
||||||
|
if self.transform:
|
||||||
|
image_tensor = self.transform(image_pil)
|
||||||
|
else:
|
||||||
|
# 如果不需要 transform,至少转为 tensor
|
||||||
|
image_tensor = torch.from_numpy(image_np).permute(2,0,1)
|
||||||
|
#image_torch = torch.from_numpy(image_np).permute(2, 0, 1) # (3,224,224)
|
||||||
|
#label_torch = torch.tensor(label_np, dtype=torch.long)
|
||||||
|
id_torch = idd
|
||||||
|
yield image_tensor,id_torch
|
||||||
|
tfrecord_path = '/kaggle/input/competitions/tpu-getting-started/tfrecords-jpeg-224x224/test/*'
|
||||||
|
dataset3 = TFRecordToPyTorchTest(tfrecord_path,transform)
|
||||||
|
test_dataloader = DataLoader(dataset3, batch_size=32, num_workers=0)
|
||||||
|
id_array=[]
|
||||||
|
all_preds=[]
|
||||||
|
with torch.no_grad():
|
||||||
|
for batch in test_dataloader:
|
||||||
|
input_ids = batch[0].to(device)
|
||||||
|
idd = batch[1]
|
||||||
|
outputs = pretrained_net(input_ids)
|
||||||
|
preds = torch.argmax(outputs, dim=1)
|
||||||
|
all_preds.extend(preds.cpu().numpy())
|
||||||
|
id_array.extend(idd)
|
||||||
|
submission = pd.DataFrame({
|
||||||
|
'id':id_array,
|
||||||
|
'label': all_preds
|
||||||
|
})
|
||||||
|
|
||||||
|
# %%
|
||||||
|
submission['id'] = submission['id'].apply(lambda x: x.decode('utf-8'))
|
||||||
|
print(submission)
|
||||||
|
submission.to_csv('submission.csv', index=False)
|
||||||
|
print("Submission saved!")
|
||||||
|
# %%
|
||||||
File diff suppressed because one or more lines are too long
Reference in New Issue
Block a user