一、项目背景

        随着医学影像技术的飞速发展,计算机断层扫描(Computed Tomography, CT)已成为临床疾病筛查、诊断与随访的核心手段之一。现代CT设备每日常规产生海量高分辨率三维图像数据,放射科医生需在短时间内对这些复杂影像进行判读,并撰写包含“影像描述”与“诊断意见”两部分的专业报告。其中,“医生描述”(即诊断意见)作为报告的核心结论,直接指导后续治疗决策,其准确性与及时性对患者预后具有决定性意义。

      在此背景下,人工智能技术,特别是多模态大模型(Multimodal Large Language Models, MLLMs)的突破,为实现从CT原始扫描数据到结构化诊断意见的端到端自动生成提供了全新可能。区别于传统仅生成客观“影像描述”的辅助工具,本项目聚焦更高阶的临床推理任务——直接基于CT原始体数据,结合患者基本信息,智能生成符合临床规范、具备循证依据的“医生描述”。

      本项目开发一个基于bart模型的自然语言生成系统,该系统可以根据CT描述生成医生描述,流程为数据处理、词表处理、预训练、微调、推理。

二、BART模型简介

      BART模型是一种用于文本生成任务的模型,在文本摘要方面表现突出,结合了bert的双向编码能力和自回归解码能力的transformer架构,在理解上下午和生成连贯摘要方面效果很好。

      编码器部分与Bert类型相似,采用双向多层Transformer架构。双向性意味着模型在处理输入文本的时候,能够同时考虑每个词的前后文信息。通过自注意力机制捕捉文本中词与词之间的依赖关系。使BART在语义编码上表现优异,为后续生成任务提供了坚实的语义基础。  

      解码器部分采用自回归机制,采用从左到右的生成方式。在生成过程中,解码器每次只根据前文信息预测一个词,是串行的。在输入的最前端加入一个start,在输出的最末端加入一个end。

三、数据预处理

      存在两个文件train.csv和test.csv保存训练数据集和测试数据集。数据包含三列,第一列为id,第二列为CT描述,第三列为医生描述。

#处理数据的文件

import pandas as pd  #处理表格数据

pre_train_file= "data/train.csv"#预处理词表

train_df = pd.read_csv(pre_train_file,header=None,names=["id","input","tgt"]) #读入数据

print(train_df.head())

# 将原始训练集划分训练集和验证集
# 使用sample方法,frac采样比例、random_state随机种子、axis轴
train_data = train_df.sample(frac=0.9, random_state=0, axis=0)   #采样0.9的比例
val_data = train_df[~train_df.index.isin(train_data.index)] #train_data.index是取到的数据的下标,train_df.index是全部数据的下标,isin包含,~取反。可以看到结果有2000条数据

# 将数据集保存为csv文件
train_data.to_csv("data/pro_train_data.csv", index=False,header=False)
val_data.to_csv("data/pro_val_data.csv", index=False,header=False)

      将原始训练集随机划分90%作为训练集,10%作为验证集保存到pro_train_data.csv和pro_val_data.csv。index=False表示不写入行索引(0,1,2),header=False表示不写入列名(表头)

四、词表处理

#处理字典的文件
import sys
import torch
from collections import Counter #collections计数工具
from transformers import BertTokenizer
from transformers import BartConfig
from transformers import BartForConditionalGeneration
from model_utils.config import parse_args

args = parse_args()         #设置 ,字典, 属性类  config  {}

# 1、读数据
def load_data(path):
    # 打开数据文件,20000行数据
    with open(path, 'r', encoding='utf-8') as f:
        lines = f.readlines()
    datas = []
    # 取出一行,每行两个","
    for line in lines:
        line = line.strip().split(",") #line.strip()去掉换行符,.split(",")按逗号分隔
        if len(line) == 3:
            # 训练集
            text, target = line[1].split(" "), line[2].split(" ") #line[?].split(" ")进一步拆分,逐字提取
            datas.append(text + target) #提取出来后加入datas中
        else:
            text = line[1].split(" ")
            datas.append(text)
    return datas

train_data = load_data('./data/train.csv')

# 2、统计所有出现过的数字
token2count = Counter()     #计数工具 哈希表
for i in train_data:
    token2count.update(i)       #不需要知道原理(调试看见统计了不重复出现的数字个数以及每个数字出现了多少次)

# 把数字从count中取出来变成列表
tail = []
ct = 0 #阈值
for k, v in token2count.items():
    if v >= ct: #超过阈值就加入列表
        tail.append(k)
tail.sort()
vocab = tail

# 3、处理词表:建立自己的词表
vocab.insert(0,"[PAD]")
vocab.insert(100,"[UNK]")
vocab.insert(101,"[CLS]")
vocab.insert(102,"[SEP]")
vocab.insert(103,"[MASK]")
vocab.insert(104,"[EOS]")

# 3、处理词表:在原词表中加字
# 在词表中查询,如果没有就加进去
# tokenizer = BertTokenizer.from_pretrained(args.pre_model_path)
# vocabs = tokenizer.get_vocab()   #获取模型词表
# print(len(vocabs))
# # 建立新词表
# new_vocabs = list(vocabs.keys())
# count = 0
# for v in vocab:         #mn复杂度
#     if v not in vocabs:
#         count += 1
#         new_vocabs.append(v)
#
# print(len(new_vocabs))
new_vocabs = vocab

# 4、保存新的词表
with open(args.pre_model_path+'/vocab.txt', 'w', encoding='utf-8') as f: #词表在mybart_base_chinese下面
    for v in new_vocabs:
        f.write(f"{v}\n")    #保存

# 4、模型部分:为什么词表变了,模型就要变
model = BartForConditionalGeneration.from_pretrained(args.pre_model_path) #原模型Embedding(1297, 768),表示词汇表大小为 1297,lm_head(768, 1297)
model.resize_token_embeddings(len(new_vocabs)) #调整模型。新模型Embedding(51440, 768),表示词汇表大小为 51440,lm_head(768, 51440)
state_dict = model.state_dict()
torch.save(state_dict, args.pre_model_path+'/pytorch_model.bin') #保存新模型
bartconfig = BartConfig.from_pretrained(args.pre_model_path) #保存config,为什么也要:因为有词表的长度设置
bartconfig.vocab_size = len(new_vocabs)
bartconfig.save_pretrained(args.pre_model_path) #新config的vocab_size从1297变成了51440

     读入train.csv数据,将一个词的ct描述和医生描述对应关系取出。并统计词出现的次数,若词出现的次数达到标准则加入词表,另外还要加入特殊字符,采用的是构造新词表。最终根据得到的词表格式调整模型的格式。

处理词表的三种方式:

(1)直接数字当id:将数据集中出现的每个数字直接作为id。但可能破坏原来该数字对应关系。(2)直接加字:将数字作为新的词加入到词表之中。但这样可能导致无用的数据在词表中,导致词表过大。

(3)构造新词表:只将数据集的数字加入到词表中。可以减少词表的大小。

五、数据预训练

1、自监督模型

class preModel(nn.Module):
    def __init__(self, args):
        super(preModel, self).__init__()
        #bart仅适合于底层特征(语义、上下文关系、语法结构这些内部的、抽象的数字表示)提取
        # self.model = BartModel.from_pretrained(args.pre_model_path)
        #预测数据会被掩盖一部分,然后使用带mlm预测头的模型(AutoModelForMaskedLM)来对数据训练,再精确
        #AutoModelForMaskedLM是自动模型加载器。1、读取args.pre_model_path目录下的config.json,2、加载model_type字段3、自动选择匹配的MLM头的模型类
        self.model = AutoModelForMaskedLM.from_pretrained(args.pre_model_path)  #BART   transformer

        print(self.model)

    # inputs里面的三个东西在MLM_Data类(pre_data)的collate里面产生的!!
    def forward(self, inputs, tgts=None):
        #inputs_ids(其中有一部分被掩盖)、attention_mask标记哪些位置是真实词,哪些是掩盖的,labels是未被掩盖的保留其原始id,掩盖的部分用-100(损失标为0,因为这些未掩盖的部分在交叉熵下loss为0,因为不变都是输入,被掩盖的部分才是输出)
        input_ids, attention_mask, labels = inputs["input_ids"], inputs["attention_mask"], inputs["labels"]
        #对非-100的label计算loss
        outputs = self.model(input_ids=input_ids, attention_mask=attention_mask, labels=labels) #把input_ids、attention_mask和label全部传给模型,这里生成模型不需要token_type_ids(即seq_ids,因为全部都是一句话)
        return outputs.loss

2、处理预训练数据集

class MLM_Data(Dataset):
    #传入句子对列表
    def __init__(self, data, args):
        super().__init__()
        self.data=data
        self.maxLen= args.input_l-3 #规定的最大长度
        self.tk=AutoTokenizer.from_pretrained(args.pre_model_path) #自动根据args.pre_model_path下的config.json加载匹配的tokenizer
        self.spNum=len(self.tk.all_special_tokens) #特殊字符的数量
        self.tkNum=self.tk.vocab_size #词表长度

    def __len__(self):
        return len(self.data)

    def random_mask(self, text_ids): #看输入是什么?ids; 输出是什么?x和y
        input_ids, output_ids = [], []
        rands = np.random.random(len(text_ids))#生成n个[0, 1)区间内的均匀分布随机浮点数
        idx=0
        while idx<len(rands):
            if rands[idx]<0.15:#需要mask,随机选
                #当决定掩码一个位置时,不只是掩盖它自己,而是随机掩码它后面的1-3个连续词(ngram),但要避免在短文本中出错,并防止大片连续掩码
                #ngram为2,掩码当前位置加后面一个词
                ngram=np.random.choice([1,2,3], p=[0.7,0.2,0.1])#若要mask,进行x_gram mask的概率
                if ngram==3 and len(rands)<7:#太大的gram不要应用于过短文本,总共才7你就要遮住3?
                    ngram=2
                if ngram==2 and len(rands)<4:
                    ngram=1
                #L和R掩码范围
                L=idx+1
                R=idx+ngram#最终需要mask的右边界(开)
                #确保掩码的其余部分的rands也小于0.15
                while L<R and L<len(rands):
                    rands[L]=np.random.random()*0.15#强制mask
                    L+=1
                idx=R
                if idx<len(rands):
                    rands[idx]=1#禁止mask片段的下一个token被mask,防止一大片连续mask
            idx+=1
        #80%的直接替换为mask,10%的随机替换词,10%的不变
        for r, i in zip(rands, text_ids):
            if r < 0.15 * 0.8:
                input_ids.append(self.tk.mask_token_id)
                output_ids.append(i)#mask预测自己
            elif r < 0.15 * 0.9:
                input_ids.append(i)
                output_ids.append(i)#自己预测自己
            elif r < 0.15:
                input_ids.append(np.random.randint(self.spNum,self.tkNum))
                output_ids.append(i)#随机的一个词预测自己,随机词不会从特殊符号中选取,有小概率抽到自己
            else:
                input_ids.append(i)
                output_ids.append(-100)#-100表示对于的词保持原样不预测

        return input_ids, output_ids

    # 取数据了!!!
    #耗时操作在此进行,可用上多进程
    def __getitem__(self, item):
        text= self.data[item]#预处理,mask等操作
        text_ids = self.tk.convert_tokens_to_ids(text) #取到一句话先转为ids
        text_ids, out_ids = self.random_mask(text_ids) #random_mask就到了MLM对数据进行的随机遮盖
        input_ids = [self.tk.cls_token_id] + text_ids + [self.tk.sep_token_id] #input_ids要加上一些特殊符号
        # [0]就是第一句话,*(len(text_ids)+2)就是只有一句话(这句代码就是segment_embeddings的处理,到bert的ppt里面看)
        #segment embedding将同一句话全标0,如果有第二句话全1
        token_type_ids=[ 0 ]*(len(text_ids)+2)
        labels = [-100] + out_ids + [-100] #label要加上一些东西,保持格式一样,就是对应inputs的cls和sep
        assert len(input_ids)==len(token_type_ids)==len(labels)
        return {'input_ids':input_ids,'token_type_ids':token_type_ids,'labels':labels}

    # 模型的输入(preModel类的inputs参数)就是在这进行转换的
    @classmethod
    def collate(cls,batch):
        input_ids=[i['input_ids'] for i in batch]
        token_type_ids=[i['token_type_ids'] for i in batch]
        labels=[i['labels'] for i in batch]
        input_ids=paddingList(input_ids,0,returnTensor=True)
        token_type_ids=paddingList(token_type_ids,0,returnTensor=True)
        labels=paddingList(labels,-100,returnTensor=True)
        attention_mask=(input_ids!=0)
        return {'input_ids':input_ids,'token_type_ids':token_type_ids
                ,'attention_mask':attention_mask,'labels':labels}

3、预训练

##预训练(自监督)1、提升泛化能力(在海量的数据上学习,让模型成为一个通才)2、降低训练成本,一次预训练,多次微调3、
from model_utils.pre_data import PreTrainDataset, loadData, MLM_Data #自定义的读数据的函数
from torch.utils.data import DataLoader, Dataset
from model_utils.models import preModel
import logging        #日志??代替打印的作用
import os
from model_utils.config import parse_args
from model_utils.utils import setup_device, setup_seed, setup_logging, build_optimizer
import torch
import time
# os.environ['CUDA_VISIBLE_DEVICES']='0'

# 与finetine.py不同处:更注重训练过程中的细节,如每个 step 的损失和剩余时间的记录,但没有验证过程。
def train_and_validate(args):
    # 1. load data  model
    model = preModel(args) #加载预训练模型
    optimizer, scheduler = build_optimizer(args, model)     #优化器设置,学习率调整
    # model = model.to(args.device)
    use_pre = False

    # 下面不重要(预训练中断的处理,单卡多卡,,,)
    if use_pre: #如果有训练好的模型就直接从保存路径加载来用
        checkpoint = torch.load(args.pre_file, map_location='cpu')
        new_KEY = model.load_state_dict(checkpoint['model_state_dict'],strict=False) #不同1:strict=False,表示在加载模型权重时,允许模型的结构与预训练模型的结构不完全一致
    if args.device == 'cuda': #选择数据串并行训练
        if args.paral == True:
            model = torch.nn.parallel.DataParallel(model.to(args.device))
        else:
            model = model.to(args.device)
        # model = BalancedDataParallel(16, model, dim=0).to(args.device)
    # model = model.to(args.device)
    #-------ema here-----------------

    # 数据部分
    all_data = loadData(args.data_path)
    train_MLM_data = MLM_Data(all_data, args)

    train_dataloader = DataLoader(train_MLM_data, batch_size=args.batch_size, shuffle=True,collate_fn=train_MLM_data.collate) #创建了训练数据集
    # 下面三行不重要
    step = 0
    start_time = time.time()
    num_total_steps = len(train_dataloader) * args.max_epochs

    # 开始训练了!!!找前向传播和梯度回传在哪里
    for epoch in range(args.max_epochs):
        for batch in train_dataloader:
            model.train()
            loss= model(batch) #这里batch里面就是数据!!!调试可以看到是长度为4的list
            #[1,2,3,4]->[2.5,2.5,2.5,2.5]
            loss = loss.mean() #多卡训练取均值
            loss.backward() #loss回传
            optimizer.step() #优化器更新
            optimizer.zero_grad() #优化器清零
            scheduler.step() #学习率调整
            #每隔若干训练步(steps),打印一次训练日志,包括当前进度、预计剩余时间(ETA)和当前 loss 值
            step += 1
            if step % args.print_steps == 0:
                time_per_step = (time.time() - start_time) / max(1, step)
                remaining_time = time_per_step * (num_total_steps - step)
                remaining_time = time.strftime('%H:%M:%S', time.gmtime(remaining_time))
                logging.info(f"Epoch {epoch} step {step} eta {remaining_time}: loss {loss:.3f}")

        logging.info(f"VAL_Epoch {epoch} step {step}: loss {loss:.3f}")
        # 不同2:预训练不验证,并且模型经过一些轮次就保存一次,不是保存最优模型
        if epoch % 5 == 0:
            torch.save({'epoch': epoch, 'model_state_dict': model.module.state_dict()},
                       f'{args.savedmodel_path}/lr{args.learning_rate}epoch{epoch}loss{loss:.3f}pre_model.bin')

def main():
    args = parse_args()  #设置(字典),单步执行进去,可以看到所有的属性都存放在里面,需要什么复制什么即可
    setup_logging()  #日志用来记录程序运行状态
    setup_device(args) #自动检测并设置运行设备,并将设备写入args
    setup_seed(args) #设置随机种子,确保结果可复现
    os.makedirs(args.savedmodel_path, exist_ok=True) #exist_ok=True是在文件夹存在时也不报错,转到声明里面看代码,创建目录
    logging.info("Training/evaluation parameters: %s", args) #打印所有设置情况,主要是LINUX需要,因为模型训练常用Linux的服务器(没有图形界面)。args用于接收命令行的参数,不用打开vim编辑器一个一个找,直接在命令行就可以设置参数
    train_and_validate(args) #进入模型训练


if __name__ == '__main__':
    main()

六、微调

1、监督模型

class myModel(nn.Module):
    def __init__(self, args):
        super(myModel, self).__init__()
        # self.model = BartModel.from_pretrained(args.pre_model_path)
        self.model = BartForConditionalGeneration.from_pretrained(args.pre_model_path) #从预训练加载的Bart模型
        self.tokenizer = AutoTokenizer.from_pretrained(args.pre_model_path)
        self.pad_id = args.pad_id
        self.tgt_pad_id = args.tgt_pad_id
        self.max_l = args.output_l
        self.beam = args.beam
        self.length_penalty = args.length_penalty
        self.no_repeat = args.no_repeat #no_repeat是什么??
        self.device = args.device

    #生成bart的输入
    def build_bart_inputs(self, input, tgt=None): #生成mask
        input_mask = (input != self.pad_id)
        if tgt == None:
            return input_mask,None
        else:
            tgt_mask = (tgt != self.tgt_pad_id)
            return input_mask, tgt_mask

    def forward(self, inputs, tgts=None):
        input_mask, tgt_mask = self.build_bart_inputs(inputs, tgts) #
        if tgts == None: #没有target#测试路径,即生成模式(里面的架构与transformer类似),串行进行
            return self.model.generate(inputs,
                                       max_length=self.max_l,
                                       attention_mask=input_mask,
                                       min_length=2,
                                       num_beams=self.beam,
                                       length_penalty=self.length_penalty,
                                       no_repeat_ngram_size=self.no_repeat,
                                       decoder_start_token_id=102
                                       # early_stopping=True,
                                       )
        outputs = self.model(input_ids=inputs, attention_mask=input_mask,
                             decoder_input_ids=tgts, decoder_attention_mask=tgt_mask)      #训练路径,并行进行
        return outputs.logits #调试可以看到logits的结构是Tensor[2, 80, 1297],表示两个样本,每个样本长度80,对80个数据都做一次1297的分类

2、验证和测试数据集


from torch.utils.data import Dataset, DataLoader
import numpy as np
import time
import csv
import traceback
from transformers import AutoTokenizer

# class BaseDataset(Dataset):
#     def _try_getitem(self, idx):
#         raise NotImplementedError
#     def __getitem__(self, idx):
#         wait = 0.1
#         while True:
#             try:
#                 ret = self._try_getitem(idx)
#                 return ret
#             except KeyboardInterrupt:
#                 break
#             except (Exception, BaseException) as e:
#                 exstr = traceback.format_exc()
#                 print(exstr)
#                 print('read error, waiting:', wait)
#                 time.sleep(wait)
#                 wait = min(wait*2, 1000)

# 不同!!!!!!
# 以前是在init里面得到x和y,这次只是把?材料?放到了init里面,x和y是在getitem里面的得到
class TranslationDataset(Dataset): #读数据
    def __init__(self, data_file, args):
        with open(data_file, 'r') as fp:
            reader = csv.reader(fp)
            self.samples = [row for row in reader][:16] #只取16个样本数据,每一条数据里面有id、ct和医生描述

            #是否全部数据
            self.input_l = args.input_l       #输入长度
            self.output_l = args.output_l       #输出长度
            self.sos_id = args.sos_id            #开始token
            self.pad_id = args.pad_id            #pad_token,多截少补
            self.eos_id = args.eos_id            # 结束
            self.tgt_pad_id = args.tgt_pad_id       # 结束pad
            self.tk=AutoTokenizer.from_pretrained(args.pre_model_path)

    def __len__(self):
        return len(self.samples)

    def __getitem__(self, idx):

        source =[self.sos_id]+ self.tk.convert_tokens_to_ids([x for x in self.samples[idx][1].split()]) + [self.eos_id] #根据下标取样本,self.samples[idx][1].split()就是把样本中的x逐字取出变成一个列表,tk.convert_tokens_to_ids把每个数字转换成词表索引,前后加上句子始终的标志
        if len(source)<self.input_l: #x长度不够就补0
            source.extend([self.pad_id] * (self.input_l-len(source)))
        if len(self.samples[idx])<3: #样本长度小于3(即测试集,没有医生描述这一列),只读x,否则还要读y
            return np.array(source)[:self.input_l]

        target = [self.sos_id] + self.tk.convert_tokens_to_ids([x for x in self.samples[idx][2].split()]) + [self.eos_id] #根据下标取样本,self.samples[idx][1].split()就是把样本中的y逐字取出变成一个列表,tk.convert_tokens_to_ids把每个数字转换成词表索引,前后加上句子始终的标志
        if len(target)<self.output_l: #y长度不够就补0
            target.extend([self.tgt_pad_id] * (self.output_l-len(target)))
        return np.array(source)[:self.input_l], np.array(target)[:self.output_l] #调试可以看到source是长度150的列表,target是长度80的列表
        #

def create_dataloaders(args, test=False):
    if not test: #如果不是测试集,就读取训练集和验证集
        train_data_path = args.data_path+"/pro_train_data.csv"
        val_data_path = args.data_path + "/pro_val_data.csv"
        train_data = TranslationDataset(train_data_path, args)
        valid_data = TranslationDataset(val_data_path, args)

        #num_workers和drop_last是什么???
        train_loader = DataLoader(train_data, batch_size=args.batch_size, shuffle=True, num_workers=args.num_workers, drop_last=False)
        valid_loader = DataLoader(valid_data, batch_size=args.val_batch_size, shuffle=True, num_workers=args.num_workers, drop_last=False)

        return train_loader, valid_loader
    else:
        test_data_path = args.data_path + "/preliminary_a_test.csv"
        test_data = TranslationDataset(test_data_path, args)
        test_loader = DataLoader(test_data, batch_size=args.test_batch_size, shuffle=False, num_workers=args.num_workers, drop_last=False)
        return test_loader

3、微调

# -*- coding: utf-8 -*-
'''这是生成任务的微调过程???'''
import logging
import os
import time
import torch
from transformers import PretrainedBartModel
from model_utils.config import parse_args
from model_utils.data import create_dataloaders
from model_utils.models import myModel
from model_utils.score import CiderD, CE
from model_utils.utils import setup_device, setup_seed, setup_logging, build_optimizer,array2str
from torch.cuda.amp import autocast as ac
from tqdm import tqdm as tqdm

os.environ['CUDA_VISIBLE_DEVICES']='0'

# 不用完全理解,关键是哪一块在做什么就行(实际上只有model和loader是自己写的),知道了以后再用到,复制就行
def validate(model, loader, args, output_file=None, beam=1, n=-1):
    res, gts = [], {}
    tot = 0
    for (source, targets) in tqdm(loader):
        if n>0 and tot>n:
            break
        source = source.cuda() #把x放到cuda上面
        pred = model(source[:, :args. input_l]) #进行预测
        pred = pred.cpu().detach().numpy() #把预测值从 GPU 移动到 CPU,并将其转换为 NumPy 数组
        #print(pred.shape)
        for i in range(pred.shape[0]): # 把预测值数组做成字典
            # res.append({'image_id':tot, 'caption': [array2str(pred[i][2:], args)]})
            # gts[tot] = [array2str(targets[i][1:], args)]
            res.append({'image_id':tot, 'caption': [array2str(pred[i], args)]}) #字典内容是id和医生描述,array2str就是把矩阵元素都变成字符串形式,输入本来是token转化成input_ids之后生成x和y,所以要把输出转换成带空格的字符串    #单步进去看
            gts[tot] = [array2str(targets[i][1:], args)] #标签也要转换成一句话
            tot += 1
    CiderD_scorer = CiderD(df='corpus', sigma=15) #这一步就是把res和gts的描述求相似度
    cider_score, cider_scores = CiderD_scorer.compute_score(gts, res)
    return cider_score

#与pretrain.py不同处:更注重验证过程中的 cider_score,并根据验证结果保存模型,同时在训练过程中明确将数据移动到 GPU
# 训练函数!!!
def train_and_validate(args):
    # 1. load data
    train_dataloader, val_dataloader = create_dataloaders(args) #加载数据
    model = myModel(args)

    #是否使用预训练模型,如果还没有预训练就设为false
    use_pre = False #已经训练过了就设为True
    if use_pre: #加载预训练过的模型
        print('use_pre')
        checkpoint = torch.load(args.my_pre_model_path, map_location='cpu')
        new_KEY = model.load_state_dict(checkpoint['model_state_dict'],strict=True) #不同1:strict=True,表示在加载模型权重时,要求模型的结构与预训练模型的结构完全一致

    optimizer, scheduler = build_optimizer(args, model)
    model = model.to(args.device)
    #-------ema here-----------------

    #进入训练!!!
    model.train()
    #-------------------------------
    # loss, results = validate(model, val_dataloader)
    # 3. training
    step = 0
    best_score = args.best_score     #评估指标,类似分类任务里面的准确率

    # 开始训练了!!!找前向传播和梯度回传在哪里
    for epoch in range(args.max_epochs):
        for (source, targets) in tqdm(train_dataloader): #读数据
            source = source.cuda() #不同2:将输入数据移动到 GPU
            targets = targets.cuda()
            # 训练模式
            model.train()
            pred = model(source[:, :args. input_l], targets[:, :args.output_l]) #得到预测值,source[:, :args. input_l]的第一个":"是样本数,第二个":"是输入长度不能超过input_l
            loss  = CE(pred[:, :-1], targets[:, 1:]) #求loss,targets里面去掉第一个(调试可以看到每个target第一个都是101,这是之前补的,所以要去掉),pred里面去掉最后一个(因为target和pred的长度要一致,而且最后一个一般都是padding这种,所以去掉最后一个)
            loss = loss.mean() #多卡训练取均值
            loss.backward() #loss回传
            optimizer.step()
            model.zero_grad()
            scheduler.step()
            step += 1
        # 验证
        if epoch % 1 == 0: #恒成立:每一轮都要做验证
            # cider_score???,用打分的模式评判文本相似率(准确率)
            cider_score = validate(model, val_dataloader, args)
            logging.info(f"Epoch {epoch} step {step}: loss {loss:.3f}, cider_score {cider_score}")
            if cider_score >= best_score: #不同3:注重验证过程中的 cider_score,并根据验证结果保存模型
                best_score = cider_score
                torch.save({'epoch': epoch, 'model_state_dict': model.state_dict()},
                        f'{args.savedmodel_path}/model_epoch_{epoch}_cider_score_{cider_score}.bin')



def main(): #和pretrain.py的代码一模一样
    args = parse_args()
    setup_logging()
    setup_device(args) #为什么设置之后就变成cpu了???
    setup_seed(args)
    os.makedirs(args.savedmodel_path, exist_ok=True)
    logging.info("Training/evaluation parameters: %s", args)
    train_and_validate(args)


if __name__ == '__main__':
    main()

七、推理

'''这是生成任务的推测过程???(测试)'''
from tqdm import tqdm
import csv
from model_utils.utils import to_device, array2str
from model_utils.models import myModel
from model_utils.data import create_dataloaders
import torch
from model_utils.config import parse_args


def inference(args):
    test_loader = create_dataloaders(args,test=True) #创建测试集
    model = myModel(args) #加载模型
    print(args.ckpt_file)

    checkpoint = torch.load(args.ckpt_file, map_location='cpu')
    model.load_state_dict(checkpoint['model_state_dict'],strict=False)
    model.to('cuda:0')
    model.eval()#用于验证/测试,既不需要训练的场景
    #用测试数据做训练
    fp = open(args.test_output_csv, 'w', newline='')
    writer = csv.writer(fp)
    tot = 0
    for source in tqdm(test_loader):
        source = to_device(source, 'cuda:0')
        pred = model(source)
        pred = pred.cpu().numpy()
        for i in range(pred.shape[0]):#pred.shape[0]=batch_size
            writer.writerow([tot, array2str(pred[i][2:], args)]) #array2str把输出转换成带空格的字符串
            tot += 1
    fp.close()

if __name__ == '__main__':
    args = parse_args()
    inference(args)

更多推荐