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