量的高效工程實(shí)踐)
1. 從“大模型”到“小模塊”為什么需要封裝ViT作為感知損失在計(jì)算機(jī)視覺(jué)的生成任務(wù)里比如圖像超分、風(fēng)格遷移或者圖像修復(fù)我們總希望生成的結(jié)果不僅像素上接近原圖更重要的是“看起來(lái)”要像。傳統(tǒng)的L1、L2損失MSE只管像素值對(duì)不對(duì)得上但人眼對(duì)紋理、結(jié)構(gòu)和語(yǔ)義的感知遠(yuǎn)比像素點(diǎn)復(fù)雜。這時(shí)候感知損失Perceptual Loss就登場(chǎng)了。它的核心思想是利用一個(gè)在大型圖像數(shù)據(jù)集如ImageNet上預(yù)訓(xùn)練好的深度網(wǎng)絡(luò)通常是VGG提取生成圖像和真實(shí)圖像在某個(gè)中間層的特征然后計(jì)算這些特征之間的差異。這個(gè)差異就代表了它們?cè)凇案兄睂用嫔系木嚯x。那么為什么現(xiàn)在大家開(kāi)始琢磨用Vision TransformerViT來(lái)替代VGG呢這事兒得從VGG的局限性說(shuō)起。VGG是個(gè)卷積神經(jīng)網(wǎng)絡(luò)CNN它的感受野是局部的通過(guò)堆疊卷積層來(lái)逐步擴(kuò)大。這意味著VGG高層特征雖然能捕捉一些全局信息但其本質(zhì)還是基于局部卷積操作的聚合。對(duì)于一些需要更強(qiáng)全局上下文理解的任務(wù)比如生成長(zhǎng)寬比較大的圖像或者圖像中物體結(jié)構(gòu)復(fù)雜、依賴遠(yuǎn)程關(guān)系的場(chǎng)景VGG可能就有點(diǎn)力不從心了。而ViT作為Transformer在視覺(jué)領(lǐng)域的成功應(yīng)用其自注意力機(jī)制天生就是為建模全局依賴關(guān)系設(shè)計(jì)的。它把圖像打成一個(gè)個(gè)Patch然后通過(guò)注意力機(jī)制讓所有Patch之間都能直接“交流”。這使得ViT提取的特征尤其是在中間層蘊(yùn)含著豐富的全局結(jié)構(gòu)和語(yǔ)義信息。直覺(jué)上用這樣的特征來(lái)計(jì)算感知損失應(yīng)該能讓生成器學(xué)會(huì)生成在結(jié)構(gòu)上更連貫、語(yǔ)義上更合理的圖像。但是直接把一個(gè)預(yù)訓(xùn)練好的ViT大模型比如ViT-B/16, ViT-L/16拿過(guò)來(lái)當(dāng)損失函數(shù)用會(huì)遇到幾個(gè)非常實(shí)際的工程問(wèn)題模型太大計(jì)算太慢內(nèi)存吃不消。一個(gè)ViT-B/16模型就有將近9000萬(wàn)參數(shù)前向傳播一次對(duì)計(jì)算資源就是不小的負(fù)擔(dān)更別說(shuō)在訓(xùn)練生成模型時(shí)每個(gè)batch、每個(gè)iteration都要計(jì)算兩次一次生成圖一次真值圖。這會(huì)讓訓(xùn)練變得極其緩慢甚至無(wú)法進(jìn)行。所以我們面臨的核心矛盾是既想利用ViT強(qiáng)大的全局感知能力又無(wú)法承受其作為損失函數(shù)帶來(lái)的巨大計(jì)算開(kāi)銷。這就引出了本文要解決的核心問(wèn)題如何將龐大的ViT模型優(yōu)雅地封裝成一個(gè)輕量、高效、即插即用的感知損失模塊。這個(gè)模塊應(yīng)該像樂(lè)高積木一樣可以輕松嵌入任何PyTorch訓(xùn)練流程中對(duì)使用者透明同時(shí)在其內(nèi)部完成模型加載、特征提取、損失計(jì)算和梯度回傳的所有臟活累活。2. 核心設(shè)計(jì)拆解一個(gè)高效ViT感知損失模塊的要素要把ViT封裝成一個(gè)好用的感知損失不能只是簡(jiǎn)單地把模型扔進(jìn)一個(gè)類里。我們需要從功能、性能和易用性三個(gè)維度進(jìn)行系統(tǒng)性的設(shè)計(jì)。一個(gè)好的封裝應(yīng)該讓用戶感覺(jué)不到背后是一個(gè)龐然大物而只是一個(gè)簡(jiǎn)單的criterion。2.1 功能設(shè)計(jì)我們需要ViT的哪一部分一個(gè)完整的ViT模型包含Patch Embedding、Transformer Encoder Blocks和最后的Classification HeadMLP。對(duì)于感知損失我們顯然不需要那個(gè)分類頭。我們的目標(biāo)是提取中間層的特征。特征層選擇和VGG感知損失通常選用relu3_3,relu4_3等層類似我們需要決定從ViT的哪個(gè)或哪些Transformer Block之后提取特征。越淺的層如第3、6塊可能包含更多細(xì)節(jié)和紋理信息越深的層如第9、12塊則包含更高級(jí)的語(yǔ)義和結(jié)構(gòu)信息。一個(gè)常見(jiàn)的策略是多層特征融合即同時(shí)提取多個(gè)中間層的特征計(jì)算加權(quán)損失這樣可以兼顧不同尺度的感知信息。特征處理ViT Encoder輸出的特征形狀通常是[Batch, Num_Patches1, Hidden_Dim]。其中Num_Patches1里的1是那個(gè)額外的[class]token。對(duì)于感知損失我們通常丟棄[class]token只使用圖像Patch對(duì)應(yīng)的特征。此外我們可能需要將這一序列特征[B, N, D]進(jìn)行重塑或池化以匹配常見(jiàn)的損失計(jì)算形式如空間維度上的MSE。2.2 性能優(yōu)化如何讓“大象”輕盈起舞這是封裝的核心挑戰(zhàn)。我們不能讓ViT的每一次前向傳播都成為訓(xùn)練瓶頸。模型凍結(jié)這是首要且必須的步驟。感知損失網(wǎng)絡(luò)在訓(xùn)練過(guò)程中參數(shù)必須被凍結(jié)requires_gradFalse。我們只是用它作為一個(gè)固定的“特征提取器”來(lái)度量圖像之間的感知距離而不是要訓(xùn)練它。這能節(jié)省大量梯度計(jì)算和內(nèi)存?;旌暇扰c設(shè)備管理混合精度AMP利用PyTorch的自動(dòng)混合精度torch.cuda.amp.autocast在特征提取時(shí)使用torch.float16半精度可以顯著減少GPU顯存占用并加速計(jì)算而對(duì)感知損失的質(zhì)量影響微乎其微。設(shè)備放置明確管理模型和輸入數(shù)據(jù)的設(shè)備。通常將ViT損失模塊放在與生成器、判別器相同的設(shè)備上如cuda:0。封裝時(shí)需要處理好輸入數(shù)據(jù)可能在不同設(shè)備上的情況。特征緩存可選但強(qiáng)力這是一個(gè)進(jìn)階優(yōu)化技巧。在像圖像到圖像翻譯這類任務(wù)中目標(biāo)圖像Ground Truth在整個(gè)訓(xùn)練過(guò)程中是固定不變的。我們可以在初始化時(shí)就預(yù)計(jì)算好所有目標(biāo)圖像在選定ViT層的特征并緩存起來(lái)。在訓(xùn)練時(shí)只需要對(duì)生成的圖像進(jìn)行前向傳播提取特征然后與緩存的特征計(jì)算損失。這直接省去了一半的ViT前向計(jì)算提速效果立竿見(jiàn)影。封裝時(shí)需要提供一個(gè)優(yōu)雅的接口來(lái)啟用和配置這個(gè)功能。2.3 接口設(shè)計(jì)如何做到“即插即用”易用性決定了這個(gè)封裝的生命力。用戶希望像使用nn.MSELoss()一樣使用它。類繼承與標(biāo)準(zhǔn)接口繼承自torch.nn.Module并實(shí)現(xiàn)forward(pred, target)方法。這是PyTorch損失函數(shù)的標(biāo)準(zhǔn)樣式用戶毫無(wú)學(xué)習(xí)成本。靈活的初始化參數(shù)允許用戶通過(guò)參數(shù)選擇model_name: 使用的ViT變體如‘vit_base_patch16_224’。feature_layers: 一個(gè)列表指定從哪些Block后提取特征如[3, 6, 9]。weights: 對(duì)應(yīng)各層特征的損失權(quán)重如[1.0, 0.5, 0.2]。use_cached_targets: 是否啟用目標(biāo)特征緩存。normalize_features: 是否對(duì)提取的特征進(jìn)行標(biāo)準(zhǔn)化如L2歸一化這有時(shí)能提升穩(wěn)定性。自動(dòng)預(yù)處理ViT預(yù)訓(xùn)練模型通常有特定的預(yù)處理要求如 resize 到 224x224使用特定的均值和標(biāo)準(zhǔn)差進(jìn)行歸一化。封裝應(yīng)該內(nèi)部集成這些預(yù)處理步驟用戶只需輸入[0,1]范圍或[0,255]范圍的RGB圖像即可無(wú)需關(guān)心細(xì)節(jié)。3. 手把手封裝從零構(gòu)建ViTPerceptualLoss類理論說(shuō)完了我們直接上代碼。下面我將一步步構(gòu)建一個(gè)功能相對(duì)完整、考慮了性能優(yōu)化的ViTPerceptualLoss類。我們會(huì)使用timm庫(kù)一個(gè)強(qiáng)大的PyTorch圖像模型庫(kù)來(lái)方便地加載預(yù)訓(xùn)練ViT。3.1 基礎(chǔ)骨架與初始化首先定義類并完成初始化工作處理模型加載、層鉤子注冊(cè)等。import torch import torch.nn as nn import torch.nn.functional as F from typing import List, Tuple, Optional import timm class ViTPerceptualLoss(nn.Module): 一個(gè)即插即用的ViT感知損失模塊。 特征提取網(wǎng)絡(luò)被凍結(jié)支持多層級(jí)特征加權(quán)可選目標(biāo)特征緩存。 def __init__(self, model_name: str vit_base_patch16_224, feature_layers: List[int] [3, 6, 9], layer_weights: List[float] None, use_cached_targets: bool False, normalize_features: bool False, input_range: str 0-1 # 0-1 or 0-255 ): super().__init__() # 參數(shù)校驗(yàn)與設(shè)置 self.feature_layers sorted(feature_layers) # 確保順序 self.normalize normalize_features self.use_cached use_cached_targets self.input_range input_range assert input_range in [0-1, 0-255], input_range must be 0-1 or 0-255 # 處理層權(quán)重 if layer_weights is None: self.layer_weights [1.0 / len(feature_layers)] * len(feature_layers) else: assert len(layer_weights) len(feature_layers), \ layer_weights must have same length as feature_layers self.layer_weights [w / sum(layer_weights) for w in layer_weights] # 歸一化 # 加載預(yù)訓(xùn)練ViT模型并凍結(jié) print(fLoading pretrained ViT: {model_name}) self.vit timm.create_model(model_name, pretrainedTrue, num_classes0) # num_classes0 移除分類頭 self.vit.eval() # 設(shè)置為評(píng)估模式 for param in self.vit.parameters(): param.requires_grad False # 注冊(cè)鉤子以捕獲中間層特征 self.features {} self._register_hooks() # 緩存目標(biāo)特征如果需要 self.target_features_cache None # 獲取模型預(yù)處理配置來(lái)自timm self.data_config timm.data.resolve_model_data_config(self.vit) self.mean torch.tensor(self.data_config[mean]).view(1, 3, 1, 1) self.std torch.tensor(self.data_config[std]).view(1, 3, 1, 1) def _register_hooks(self): 為選定的Transformer Blocks注冊(cè)前向鉤子捕獲其輸出。 def get_feature_hook(layer_id): def hook(module, input, output): # output 通常是 tuple我們?nèi)〉谝粋€(gè)通常是經(jīng)過(guò)Block處理后的tensor # 形狀: [B, N1, D] self.features[layer_id] output[0] if isinstance(output, tuple) else output return hook # timm的ViT模型blocks通常存儲(chǔ)在 blocks 屬性中 for i, layer_idx in enumerate(self.feature_layers): layer self.vit.blocks[layer_idx] layer.register_forward_hook(get_feature_hook(layer_idx))關(guān)鍵點(diǎn)解析timm.create_model(..., num_classes0)num_classes0是關(guān)鍵它告訴timm我們不需要最后的分類頭模型直接返回最后一個(gè)Transformer Block輸出的特征。這正好符合我們的需求。self.vit.eval()和param.requires_gradFalse雙保險(xiǎn)確保模型在訓(xùn)練我們的生成器時(shí)不會(huì)被意外更新同時(shí)啟用BatchNorm/ LayerNorm的推理模式。鉤子Hook機(jī)制這是動(dòng)態(tài)獲取中間層輸出的標(biāo)準(zhǔn)方法。我們?cè)谥付ǖ腷locks[layer_idx]上注冊(cè)鉤子當(dāng)前向傳播執(zhí)行到該層時(shí)鉤子函數(shù)會(huì)被調(diào)用我們將輸出存儲(chǔ)到self.features字典中鍵就是層索引。數(shù)據(jù)配置timm為每個(gè)預(yù)訓(xùn)練模型提供了標(biāo)準(zhǔn)的預(yù)處理參數(shù)均值、標(biāo)準(zhǔn)差、輸入尺寸。我們?cè)谶@里獲取它以便在forward函數(shù)中進(jìn)行一致的預(yù)處理。3.2 核心前向傳播與損失計(jì)算接下來(lái)實(shí)現(xiàn)forward方法這是模塊的核心。def _preprocess(self, x: torch.Tensor) - torch.Tensor: 將輸入圖像預(yù)處理為ViT模型期望的格式。 # 1. 確保輸入是4D Tensor [B, C, H, W] if x.dim() 3: x x.unsqueeze(0) # 2. 調(diào)整輸入范圍到 [0, 1] if self.input_range 0-255: x x / 255.0 # 3. 調(diào)整大小到模型期望的尺寸 (例如 224x224) # 注意雙線性插值通常對(duì)感知損失影響不大因?yàn)閾p失基于特征而非像素。 target_size self.data_config[input_size][1:] # 假設(shè)是 (224, 224) if x.shape[-2:] ! target_size: x F.interpolate(x, sizetarget_size, modebilinear, align_cornersFalse) # 4. 使用模型特定的均值和標(biāo)準(zhǔn)差進(jìn)行歸一化 device x.device x (x - self.mean.to(device)) / self.std.to(device) return x def _extract_vit_features(self, x: torch.Tensor) - List[torch.Tensor]: 通過(guò)ViT網(wǎng)絡(luò)前向傳播并返回指定層的特征列表。 清空之前的特征緩存提取新特征。 self.features.clear() # 清除上一次的特征 with torch.no_grad(): # 無(wú)需梯度節(jié)省內(nèi)存 # 注意我們只運(yùn)行到足以獲取所需特征層的位置。 # 但timm模型通常需要完整前向。這里簡(jiǎn)單處理運(yùn)行整個(gè)網(wǎng)絡(luò)。 # 由于鉤子已注冊(cè)運(yùn)行時(shí)會(huì)自動(dòng)填充 self.features _ self.vit(x) # 按 self.feature_layers 的順序收集特征 extracted_features [] for layer_idx in self.feature_layers: feat self.features[layer_idx] # [B, N1, D] # 移除 [class] token只保留圖像patch特征 feat feat[:, 1:, :] # [B, N, D] # 可選對(duì)特征進(jìn)行L2歸一化 if self.normalize: feat F.normalize(feat, p2, dim-1) extracted_features.append(feat) return extracted_features def forward(self, pred: torch.Tensor, target: torch.Tensor, target_cache_id: Optional[str] None) - torch.Tensor: 計(jì)算預(yù)測(cè)圖像與目標(biāo)圖像之間的ViT感知損失。 Args: pred: 預(yù)測(cè)圖像形狀 [B, C, H, W] target: 目標(biāo)圖像形狀 [B, C, H, W] target_cache_id: 可選用于標(biāo)識(shí)和檢索緩存的目標(biāo)特征。如果為None且啟用緩存則使用默認(rèn)緩存。 Returns: 標(biāo)量損失值。 # 0. 設(shè)備同步 device pred.device self.vit.to(device) self.mean self.mean.to(device) self.std self.std.to(device) # 1. 預(yù)處理 pred_preprocessed self._preprocess(pred) target_preprocessed self._preprocess(target) # 2. 提取預(yù)測(cè)圖像的特征 pred_features_list self._extract_vit_features(pred_preprocessed) # 3. 獲取目標(biāo)圖像的特征 (可能來(lái)自緩存) if self.use_cached and self.target_features_cache is not None: # 從緩存中獲取目標(biāo)特征 if target_cache_id is not None: target_features_list self.target_features_cache[target_cache_id] else: # 使用默認(rèn)緩存假設(shè)batch size為1或已預(yù)先緩存了整個(gè)目標(biāo)集 target_features_list self.target_features_cache[default] else: # 實(shí)時(shí)提取目標(biāo)特征 with torch.no_grad(): target_features_list self._extract_vit_features(target_preprocessed) # 如果啟用緩存且是第一次則進(jìn)行緩存 if self.use_cached and self.target_features_cache is None: self.target_features_cache {default: target_features_list} # 4. 計(jì)算加權(quán)感知損失 total_loss 0.0 for w, pred_feat, target_feat in zip(self.layer_weights, pred_features_list, target_features_list): # 使用L2損失MSE或L1損失。L1有時(shí)更穩(wěn)定。 # layer_loss F.mse_loss(pred_feat, target_feat) layer_loss F.l1_loss(pred_feat, target_feat) total_loss w * layer_loss return total_loss def cache_target_features(self, target_images: torch.Tensor, cache_id: str default): 預(yù)計(jì)算并緩存一批目標(biāo)圖像的特征。 這在訓(xùn)練開(kāi)始前調(diào)用一次可以極大加速訓(xùn)練。 Args: target_images: 目標(biāo)圖像Tensor形狀 [N, C, H, W] cache_id: 緩存標(biāo)識(shí)符 if not self.use_cached: print(Warning: use_cached is False, caching will have no effect.) return device target_images.device self.vit.to(device) target_preprocessed self._preprocess(target_images) with torch.no_grad(): features self._extract_vit_features(target_preprocessed) if self.target_features_cache is None: self.target_features_cache {} self.target_features_cache[cache_id] features print(fTarget features cached for id: {cache_id})關(guān)鍵點(diǎn)解析_preprocess封裝了所有繁瑣的預(yù)處理步驟用戶無(wú)需關(guān)心。注意其中的interpolate將輸入圖像縮放到ViT的標(biāo)準(zhǔn)輸入尺寸如224x224。這是必須的因?yàn)轭A(yù)訓(xùn)練ViT的Patch Embedding是固定大小的。_extract_vit_features這是特征提取的核心。with torch.no_grad()確保了在提取特征時(shí)不會(huì)計(jì)算和存儲(chǔ)梯度節(jié)省大量顯存。feat[:, 1:, :]這行代碼去掉了[class]token因?yàn)槲覀冴P(guān)心的是圖像區(qū)域的特征。forward中的緩存邏輯這是性能優(yōu)化的關(guān)鍵。如果use_cachedTrue并且我們已經(jīng)通過(guò)cache_target_features方法預(yù)計(jì)算了目標(biāo)特征那么在訓(xùn)練循環(huán)中target圖像的特征就直接從緩存中讀取省去了對(duì)target圖像的ViT前向傳播。這對(duì)于固定目標(biāo)數(shù)據(jù)集的訓(xùn)練如超分、去噪提速效果極其顯著。損失函數(shù)選擇代碼中使用了F.l1_loss。在感知損失中L1損失MAE通常比L2損失MSE更魯棒因?yàn)樗鼘?duì)異常值不那么敏感能產(chǎn)生更清晰的圖像。這是一個(gè)經(jīng)驗(yàn)性的選擇。3.3 在訓(xùn)練循環(huán)中使用封裝好后使用起來(lái)就非常簡(jiǎn)單了。# 1. 初始化損失函數(shù) perceptual_loss_fn ViTPerceptualLoss( model_namevit_base_patch16_224, feature_layers[3, 6, 9], layer_weights[1.0, 0.8, 0.5], use_cached_targetsTrue, # 啟用緩存 normalize_featuresTrue, input_range0-1 ).cuda() # 2. 可選但推薦如果目標(biāo)數(shù)據(jù)集是固定的如訓(xùn)練集預(yù)緩存特征 # 假設(shè) train_target_loader 是加載目標(biāo)圖像的DataLoader all_targets [] for target_batch in train_target_loader: all_targets.append(target_batch.cuda()) all_targets torch.cat(all_targets, dim0) perceptual_loss_fn.cache_target_features(all_targets, cache_idtrain_set) # 3. 在訓(xùn)練循環(huán)中 for epoch in range(num_epochs): for batch_idx, (input_imgs, target_imgs) in enumerate(train_loader): input_imgs, target_imgs input_imgs.cuda(), target_imgs.cuda() # 生成圖像 generated_imgs generator(input_imgs) # 計(jì)算損失 mse_loss F.mse_loss(generated_imgs, target_imgs) # 使用緩存?zhèn)魅雝arget_imgs主要是為了形狀匹配實(shí)際特征從緩存中按索引或批次獲取。 # 這里假設(shè)DataLoader順序固定可以使用batch_idx或其他ID。更穩(wěn)健的做法是使用圖像本身的ID。 # 簡(jiǎn)化示例我們假設(shè)緩存了所有目標(biāo)且順序一致這里直接使用默認(rèn)緩存。 perc_loss perceptual_loss_fn(generated_imgs, target_imgs) # target_imgs在啟用緩存時(shí)僅用于占位和獲取設(shè)備信息 total_loss mse_loss 0.1 * perc_loss # 加權(quán)總和 optimizer.zero_grad() total_loss.backward() optimizer.step()4. 高級(jí)技巧、避坑指南與效果對(duì)比把模塊跑起來(lái)只是第一步要想讓它真正發(fā)揮作用還需要一些細(xì)節(jié)上的打磨和對(duì)潛在問(wèn)題的預(yù)判。4.1 特征層與權(quán)重的調(diào)參經(jīng)驗(yàn)選擇哪些層以及賦予多大權(quán)重是影響感知損失效果的關(guān)鍵。淺層如第1-4塊更多地捕捉邊緣、紋理、顏色等低級(jí)特征。如果你的任務(wù)側(cè)重于紋理合成或細(xì)節(jié)恢復(fù)如紋理超分可以賦予淺層更高的權(quán)重。中層如第5-8塊開(kāi)始捕捉更復(fù)雜的圖案和部件信息。這是一個(gè)比較平衡的選擇適用于大多數(shù)通用圖像生成任務(wù)。深層如第9-12塊捕捉高級(jí)語(yǔ)義和全局結(jié)構(gòu)。如果你的任務(wù)對(duì)物體的形狀和布局要求很高如語(yǔ)義分割圖生成照片深層特征就尤為重要。實(shí)戰(zhàn)建議從[3, 6, 9]這樣的均勻分布開(kāi)始嘗試權(quán)重設(shè)為[1.0, 1.0, 1.0]。然后根據(jù)生成結(jié)果調(diào)整。如果發(fā)現(xiàn)結(jié)果過(guò)于平滑、缺乏細(xì)節(jié)就增加淺層權(quán)重如果發(fā)現(xiàn)結(jié)構(gòu)扭曲就增加深層權(quán)重。一個(gè)常見(jiàn)的策略是使用所有層但給深層一個(gè)衰減的權(quán)重例如list(range(12))配合[1.0]*12的權(quán)重或者指數(shù)衰減的權(quán)重。4.2 內(nèi)存與速度的終極優(yōu)化梯度檢查點(diǎn)與特征蒸餾即使凍結(jié)了ViT前向傳播的內(nèi)存占用對(duì)于大batch size或高分辨率圖像需要插值到224依然可能是個(gè)問(wèn)題。梯度檢查點(diǎn)Gradient Checkpointing這是用計(jì)算時(shí)間換顯存的神器。PyTorch的torch.utils.checkpoint可以讓我們只保存部分中間結(jié)果在反向傳播時(shí)重新計(jì)算其余部分。對(duì)于ViT這種多層Transformer可以對(duì)其中的某些Block應(yīng)用檢查點(diǎn)。但是請(qǐng)注意我們的ViT是凍結(jié)的不需要反向傳播梯度給它的參數(shù)。因此標(biāo)準(zhǔn)的梯度檢查點(diǎn)在這里不適用。我們主要需要節(jié)省的是前向傳播的**激活值A(chǔ)ctivations**占用的顯存。一個(gè)變通的方法是在_extract_vit_features方法中用torch.no_grad()包裹整個(gè)前向這樣PyTorch就不會(huì)保存中間激活值用于反向傳播因?yàn)楦静恍枰獜亩烊还?jié)省了這部分顯存。我們代碼中已經(jīng)這么做了。特征蒸餾訓(xùn)練一個(gè)輕量化的“代理”網(wǎng)絡(luò)如果ViT的計(jì)算成本在您的場(chǎng)景下仍然無(wú)法接受終極方案是知識(shí)蒸餾。你可以先用完整的ViT感知損失在一個(gè)小型數(shù)據(jù)集上訓(xùn)練你的生成器。同時(shí)訓(xùn)練一個(gè)輕量級(jí)的CNN如一個(gè)小型ResNet或MobileNet讓它去學(xué)習(xí)模仿ViT中間層的特征輸出。訓(xùn)練完成后用這個(gè)輕量級(jí)CNN替代ViT作為感知損失。這樣你既保留了ViT強(qiáng)大的感知能力又獲得了CNN的推理速度。這需要額外的訓(xùn)練步驟但是一次投入長(zhǎng)期受益。4.3 常見(jiàn)坑點(diǎn)與排查清單輸入范圍錯(cuò)誤這是最常見(jiàn)的錯(cuò)誤。預(yù)訓(xùn)練ViT期望的輸入是經(jīng)過(guò)特定均值和標(biāo)準(zhǔn)差歸一化的。我們的_preprocess方法封裝了它。請(qǐng)務(wù)必確認(rèn)你傳入的圖像Tensor范圍是[0,1]還是[0,255]并通過(guò)input_range參數(shù)正確設(shè)置。特征形狀不匹配當(dāng)你嘗試計(jì)算F.l1_loss(pred_feat, target_feat)時(shí)確保兩個(gè)特征張量形狀完全一致。如果啟用了緩存要確保緩存的target_feat和當(dāng)前pred_feat的batch size能對(duì)應(yīng)上或者通過(guò)廣播機(jī)制兼容。在緩存時(shí)最好緩存整個(gè)數(shù)據(jù)集的特征然后在forward中根據(jù)索引來(lái)取對(duì)應(yīng)的特征批次。損失值為NaN或爆炸首先檢查輸入圖像是否有異常值如超出范圍。其次嘗試對(duì)特征進(jìn)行L2歸一化normalize_featuresTrue這能穩(wěn)定訓(xùn)練。最后可以降低感知損失的權(quán)重如從0.1降到0.01或0.001因?yàn)樗赡苤鲗?dǎo)了梯度。緩存導(dǎo)致的數(shù)據(jù)泄露在類似圖像翻譯的任務(wù)中如果訓(xùn)練集和驗(yàn)證集的目標(biāo)圖像不同務(wù)必為它們創(chuàng)建不同的緩存ID如cache_target_features(..., ‘train’)和cache_target_features(..., ‘val’)并在驗(yàn)證時(shí)使用對(duì)應(yīng)的ID。切忌在驗(yàn)證時(shí)錯(cuò)誤地使用訓(xùn)練集的緩存特征。ViT模型選擇timm提供了眾多ViT變體。vit_base_patch16_224是一個(gè)不錯(cuò)的起點(diǎn)。如果你想減少計(jì)算量可以嘗試vit_small_patch16_224或vit_tiny_patch16_224。但請(qǐng)注意模型越小其感知能力可能越弱需要權(quán)衡。4.4 與VGG感知損失的直觀對(duì)比為了讓你有個(gè)直觀感受我簡(jiǎn)單對(duì)比一下在同一個(gè)圖像著色任務(wù)上使用VGG-19relu3_3和ViT-B/16[3,6,9]層作為感知損失的效果差異基于個(gè)人實(shí)驗(yàn)經(jīng)驗(yàn)細(xì)節(jié)與紋理VGG損失傾向于生成紋理更豐富、細(xì)節(jié)更銳利的結(jié)果但有時(shí)會(huì)顯得有點(diǎn)“碎”或過(guò)度紋理化。ViT損失生成的紋理更自然、連貫尤其是在有重復(fù)模式或長(zhǎng)程結(jié)構(gòu)如建筑立面、森林的場(chǎng)景中。全局結(jié)構(gòu)與一致性ViT損失在維持圖像全局結(jié)構(gòu)一致性上表現(xiàn)明顯更好。例如在生成長(zhǎng)線條如地平線、建筑輪廓時(shí)ViT損失引導(dǎo)的結(jié)果線條更直扭曲更少。VGG由于感受野限制有時(shí)會(huì)導(dǎo)致長(zhǎng)距離結(jié)構(gòu)出現(xiàn)彎曲或不連續(xù)。語(yǔ)義合理性對(duì)于需要高級(jí)語(yǔ)義理解的任務(wù)比如根據(jù)草圖生成物體ViT損失能更好地避免語(yǔ)義錯(cuò)誤比如把貓的耳朵生成在錯(cuò)誤的位置。計(jì)算成本毫無(wú)疑問(wèn)VGG-19的計(jì)算速度遠(yuǎn)快于ViT-B/16。即使經(jīng)過(guò)我們的優(yōu)化凍結(jié)、緩存ViT損失的計(jì)算開(kāi)銷仍然是VGG的數(shù)倍。所以選擇哪一個(gè)如果你的任務(wù)對(duì)細(xì)節(jié)紋理要求極高且計(jì)算資源有限VGG感知損失依然是可靠的選擇。如果你的任務(wù)強(qiáng)調(diào)整體結(jié)構(gòu)、長(zhǎng)程依賴和語(yǔ)義正確性并且你有一定的GPU算力那么封裝好的ViT感知損失會(huì)帶來(lái)質(zhì)的提升。