Pytorch中項(xiàng)目配置文件的管理與導(dǎo)入方式
1.yaml文件
在 PyTorch 深度學(xué)習(xí)項(xiàng)目中,使用 YAML(Yet Another Markup Language)作為配置文件是非常主流的做法。相比 JSON 或 XML,YAML 的可讀性更強(qiáng),非常適合用來管理復(fù)雜的超參數(shù)(Hyperparameters)、模型結(jié)構(gòu)參數(shù)和文件路徑。
1.1 為什么是yaml
在深度學(xué)習(xí)中,我們經(jīng)常需要調(diào)整 batch_size, learning_rate, optimizer 等參數(shù)。
- 如果不使用配置: 你需要反復(fù)修改代碼中的變量,容易出錯(cuò)且難以版本控制。
- 使用 YAML: 將代碼(邏輯)與參數(shù)(配置)分離。修改參數(shù)只需改動 YAML 文件,無需觸碰核心代碼。
1.2 文件的編寫語法
YAML 的核心規(guī)則是依靠縮進(jìn)(Indentation)來表示層級關(guān)系。基本的語法概括如下:
- 縮進(jìn): 必須使用空格,不能使用 Tab 鍵(通常是 2 個(gè)或 4 個(gè)空格)。
- 鍵值對:
key: value(冒號后面必須有一個(gè)空格)。 - 注釋: 使用
#。
細(xì)致總結(jié)一下:
1.大小寫敏感:True 和 true 是不同的(YAML 對“布爾值的關(guān)鍵字”不區(qū)分大小寫,但 YAML 對“字符串內(nèi)容”是區(qū)分大小寫的)。
2.縮進(jìn)表示層級關(guān)系:
- 只能使用空格(Space),絕對不能用 Tab 鍵(這是 YAML 最常見的錯(cuò)誤來源)。
- 縮進(jìn)空格數(shù)不固定(可以是 2 個(gè)或 4 個(gè)),但同一層級必須對齊,子層級必須比父層級多縮進(jìn)。
- 示例(你的文件用 2 個(gè)空格):
paths: # 第 0 級 data_dir: "./data/cifar10" # 第 1 級(縮進(jìn) 2 空格) log_dir: "./logs/experiment_1" #如果縮進(jìn)不一致(如一個(gè) 2 空格、一個(gè) 4 空格),解析器會報(bào)錯(cuò)。
3.鍵值對:格式為 key: value(冒號后必須有一個(gè)空格)。如果沒空格,如 key:value,會解析失敗。
4.注釋:用 # 開頭,從 #到行尾都被忽略??梢苑旁谛惺住⑿形不騿为?dú)一行。示例:
use_gpu: true # 布爾值(注釋在行尾) # 路徑配置(單獨(dú)一行注釋) paths: ...
5.文檔分隔:一個(gè)文件中可以有多個(gè) YAML 文檔,用 — 分隔。例如
文檔分隔的作用:
邏輯上將一個(gè)文件拆分成多個(gè)獨(dú)立的配置對象:每個(gè) — 之前的部分是一個(gè)完整的、獨(dú)立的 YAML 文檔(相當(dāng)于一個(gè)獨(dú)立的字典、配置或數(shù)據(jù)結(jié)構(gòu))。
允許在同一個(gè)文件中存儲多個(gè)相關(guān)或不相關(guān)的配置,而不需要拆分成多個(gè)物理文件。
方便某些工具一次性處理多個(gè)配置,比如批量導(dǎo)入、流水線處理等。
#YAML 文件的標(biāo)準(zhǔn)規(guī)范允許一個(gè)物理文件中包含多個(gè)獨(dú)立的 YAML 文檔(相當(dāng)于多個(gè)獨(dú)立的配置對象),它們之間用 ---(三個(gè)連字符)來分隔。 # 第一個(gè) YAML 文檔 name: Alice age: 30 hobbies: - reading - hiking --- #用三個(gè)連字符或者三個(gè)點(diǎn)···來顯示結(jié)束一個(gè)文檔,通常不需要。如果在yaml文件中如果有和yaml沒關(guān)系的內(nèi)容,必須有結(jié)束符號。 # 第二個(gè) YAML 文檔 name: Bob age: 25 hobbies: - gaming - cooking --- # 第三個(gè) YAML 文檔 server: host: localhost port: 8080
6.數(shù)據(jù)類型詳解–YAML 支持三種基本結(jié)構(gòu):
- 標(biāo)量(Scalars):單個(gè)值(如字符串、數(shù)字、布爾)。
- 映射(Mappings):鍵值對集合(相當(dāng)于字典/dict)。
- 序列(Sequences):有序列表(相當(dāng)于數(shù)組/list)。
(1) 字符串(String)最常見類型。**YAML 里的“字符串”,本質(zhì)就是:一段文字。不同寫法,只是“YAML 怎么把這段文字當(dāng)成什么樣子來理解”。**可以不加引號(plain style):如果不含特殊字符(如 : { } [ ] , #),推薦不加引號,更簡潔。
- 示例:data_dir: “./data/cifar10”(路徑通常加引號,避免解析問題)。路徑里可能有特殊字符,YAML 解析器容易誤解
- 單引號 ‘…’:內(nèi)容原樣輸出,雙單引號 ‘’ 表示單個(gè) '。
單引號 '...'(原樣保存)
msg: 'hello\nworld' #不是換行實(shí)際上結(jié)果是 "hello\\nworld"------Python 用 \\ 來表示“字符串里有一個(gè)反斜杠” # 單引號 `'...'`(原樣保存) 不做任何的轉(zhuǎn)義。
當(dāng)連續(xù)出現(xiàn)兩個(gè)單引號的時(shí)候 ‘’ 表示單個(gè)單引號;
msg: 'it''s good' #等價(jià)于 "it's good"
雙引號 “…”:支持轉(zhuǎn)義(如 \n 換行、\t Tab)。支持轉(zhuǎn)義符(和 Python 字符串一樣)
msg: "hello\nworld" #雙引號支持轉(zhuǎn)義所以結(jié)果是 "hello world"
多行字符串:
| :保留換行(literal block)。| = “我寫幾行,你就給我?guī)仔?rdquo;
description: | this is line one this is line two this is line three #在python里邊 "this is line one\nthis is line two\nthis is line three\n"
>:折疊換行成空格(folded block)。> = “寫的時(shí)候換行,讀的時(shí)候當(dāng)一行”
description: > this is line one this is line two this is line three #實(shí)際上 "this is line one this is line two this is line three\n"
(2)數(shù)字在深度學(xué)習(xí) YAML 里,大多數(shù)只需要會這三種數(shù)字寫法:
epochs: 100 # 整數(shù) lr: 0.001 # 小數(shù) weight_decay: 1.0e-4 # 科學(xué)計(jì)數(shù)法----1e-4 = 1 × 10?? = 0.0001 #其他格式:十六進(jìn)制 0xFF、八進(jìn)制 0o777。
(3)布爾值標(biāo)準(zhǔn)寫法:true / false(小寫推薦)。YAML 也支持變體:True、TRUE、Yes、No、On、Off(不區(qū)分大小寫)。但是注意,不能加引號,不然會變成字符串
(4) Null(空值): 用 ~ 或 null 表示。
#例如 optional: ~
(5)映射(字典/Dict):YAML 的“映射(Mapping)”= Python 的“字典(dict)”, 本質(zhì)就是:鍵 → 值 的對應(yīng)關(guān)系
#縮進(jìn)只能用空格,不能用 Tab # 對 train: batch_size: 64 # 錯(cuò)(Tab) train: ?batch_size: 64 #不能用tab鍵 #同一層級,縮進(jìn)必須對齊 # 對 train: batch_size: 64 epochs: 100 # 錯(cuò) train: batch_size: 64 epochs: 100 #必須唯一(同一層里) train: batch_size: 64 batch_size: 128 # 覆蓋 / 非法 #冒號分左右,縮進(jìn)分里外,對齊是同級,一切都是鍵值對 #YAML 的映射不是“復(fù)雜”,而是“把 Python dict 寫得更好看”
(6)序列:序列 = 一堆有順序的元素,類似于python里邊的list。
#block風(fēng)格 transform_list: #transform_list: → 一個(gè)鍵 - "RandomCrop" # - → 一個(gè)列表元素,每個(gè) - 表示一項(xiàng) - "RandomHorizontalFlip" - "Normalize" #看到 -,就要想到“列表的一項(xiàng)” ,‘-‘后邊一定要有空格 #flow風(fēng)格,行內(nèi)寫法 transform_list: ["RandomCrop", "RandomHorizontalFlip", "Normalize"] #一般在:列表很短,不嵌套,不需要注釋。時(shí)使用
YAML = 映射(dict) + 序列(list) + 標(biāo)量(string / number / bool)
# config.yaml project_name: "ResNet_Classification" use_gpu: true # 布爾值 # 路徑配置 paths: data_dir: "./data/cifar10" log_dir: "./logs/experiment_1" # 模型參數(shù) model: type: "resnet18" num_classes: 10 pretrained: true # 訓(xùn)練超參數(shù) train: batch_size: 64 epochs: 100 learning_rate: 0.001 weight_decay: 1.0e-4 # 支持科學(xué)計(jì)數(shù)法 optimizer: "Adam" # 列表/數(shù)組寫法 transform_list: - "RandomCrop" - "RandomHorizontalFlip" - "Normalize"
1.3 yaml文件的使用
1.使用 yaml
import yaml
# 讀取函數(shù)
def get_config(path):
with open(path, 'r', encoding='utf-8') as f:
return yaml.safe_load(f) #把yaml文件內(nèi)容轉(zhuǎn)換為python字典。safe_load 推薦用,不會執(zhí)行 YAML 文件里潛在的危險(xiǎn)命令。
cfg = get_config("config.yaml") #輸入路徑
# 使用方式:像查字典一樣
print(cfg['learning_rate']) # 輸 出結(jié)果
# 缺點(diǎn):如果層級很深,代碼會變成 config['train']['params']['lr'],很難看且容易寫錯(cuò)字符串
2.封裝為對象
在日常的項(xiàng)目中,我們不希望在代碼里寫滿 ['key']。我們更習(xí)慣用 . 來訪問屬性,比如 config.lr。
#利用SimpleNamespace實(shí)現(xiàn)---SimpleNamespace 是 Python 標(biāo)準(zhǔn)庫(types 模塊)中的一個(gè)非常輕量的類。它的作用:允許你動態(tài)地給一個(gè)對象添加屬性,并用點(diǎn)號訪問這些屬性。相當(dāng)于一個(gè)“可隨意擴(kuò)展屬性的空對象”。
import yaml #yaml 是 import 的 PyYAML 庫。
from types import SimpleNamespace
def load_config_as_obj(yaml_path):
"""
讀取 yaml 并將字典遞歸轉(zhuǎn)換為對象,方便用 . 屬性訪問
"""
with open(yaml_path, 'r', encoding='utf-8') as f: #open打開文件,返回一個(gè)文件對象f
config_dict = yaml.safe_load(f) #加載為python字典
# 遞歸轉(zhuǎn)換函數(shù)
def dict_to_obj(d):
if not isinstance(d, dict):
'''
#isinstance(object, class_or_tuple):--判斷這個(gè)對象是不是某種類型
object:你要檢查的變量.
class_or_tuple:你想檢查的類型(或者類型元組)
返回值:布爾值 True / False
'''
return d
# 將字典轉(zhuǎn)為 SimpleNamespace 對象
obj = SimpleNamespace() #創(chuàng)建一個(gè)空的SimpleNamespace對象。調(diào)用 types.SimpleNamespace 類,創(chuàng)建一個(gè)空的、可動態(tài)加屬性的對象。此時(shí) obj 里面什么屬性都沒有。
for k, v in d.items(): #d.items返回的是一個(gè)元組,for循環(huán)可以多個(gè)變量,但是要求可迭代對象的每個(gè)元素是元組或列表,元素的長度必須和變量數(shù)一致
# 遞歸處理嵌套的字典
setattr(obj, k, dict_to_obj(v)) #這里遞歸調(diào)用dict_to_obj函數(shù)。如果不是字典,則返回d(也就是v)。如果是字典在進(jìn)來再進(jìn)行調(diào)用,直到不是字典未知。----給對象 obj 動態(tài)增加一個(gè)屬性,名字是 k,值是 dict_to_obj(v)。
#setattr(object, name, value)---把 name 當(dāng)作屬性名,把 value 賦值給對象
'''
object:要操作的對象
name:屬性名(字符串)
value:要賦給屬性的值
'''
return obj #可調(diào)用對象
return dict_to_obj(config_dict) #返回值
# --- 使用演示 ---
# 假設(shè) yaml 內(nèi)容是:
# train:
# lr: 0.01
# device: "cuda"
cfg = load_config_as_obj("config.yaml") #給一個(gè)yaml文件路徑
# 現(xiàn)在的調(diào)用方式非常優(yōu)雅:
print(cfg.train.lr) # 0.01
print(cfg.train.device) # cuda
其次可以使用 argparse 讀取命令行參數(shù),如果有輸入,就覆蓋 yaml 里的默認(rèn)值。
argparse 是 Python 內(nèi)置模塊,用來 解析命令行參數(shù)。“命令行參數(shù)” = 你運(yùn)行腳本時(shí)輸入的參數(shù),比如:
python train.py --lr 0.001 --epochs 50
argparse 可以把這些字符串參數(shù) 轉(zhuǎn)換成 Python 對象,方便在代碼中使用使用 argparse 通常有三個(gè)步驟:
創(chuàng)建解析器
parser = argparse.ArgumentParser()
ArgumentParser() 創(chuàng)建一個(gè)解析器對象。這個(gè)解析器負(fù)責(zé)定義你想接受哪些參數(shù),以及解析命令行輸入
添加參數(shù)定義:
add_argument 是 argparse 模塊里 ArgumentParser 對象的方法,作用是:告訴解析器你的程序可以接收哪些命令行參數(shù),以及這些參數(shù)的類型、默認(rèn)值和說明。
parser.add_argument('--lr', type=float, default=None, help='學(xué)習(xí)率')
--lr→ 命令行參數(shù)名type=float→ 解析后轉(zhuǎn)換為浮點(diǎn)數(shù)default=None→ 如果命令行沒提供,默認(rèn)值是 Nonehelp='學(xué)習(xí)率'→ 提示信息(python train.py --help會顯示)
你可以添加多個(gè)參數(shù):
parser.add_argument('--epochs', type=int, default=None, help='訓(xùn)練輪數(shù)')
parser.add_argument('--config', type=str, default='./configs/resnet_train.yaml', help='配置文件路徑')
解析命令行輸入
args = parser.parse_args() #`parse_args()` 會讀取運(yùn)行腳本時(shí)的命令行參數(shù),返回一個(gè)對象 `args`,里面每個(gè)參數(shù)都是 **對象屬性**
例如運(yùn)行:
python train.py --lr 0.001 --epochs 50
得到:
args.lr # 0.001 args.epochs # 50 args.config # './configs/resnet_train.yaml' #如果命令行不輸入某個(gè)參數(shù),它就用你定義的 `default` 值。
import argparse #Python 內(nèi)置模塊,用來解析命令行參數(shù)(命令行參數(shù)也就是python運(yùn)行腳本的時(shí)候輸入的參數(shù):python train.py --lr 0.001 --epochs 50)。argparse 可以把這些字符串參數(shù)轉(zhuǎn)換成 Python 對象,方便在代碼中使用
def get_args_and_config(): #讀取 YAML 配置 + 解析命令行參數(shù) + 覆蓋默認(rèn)值。返回最終的 cfg 對象,用于訓(xùn)練腳本中直接訪問參數(shù)
parser = argparse.ArgumentParser() #ArgumentParser() 創(chuàng)建一個(gè)解析器對象,知道你程序允許哪些命令行參數(shù),并解析這些參數(shù)
parser.add_argument('--config', type=str, default='./configs/resnet_train.yaml', help='配置文件路徑') #拿到y(tǒng)aml文件的路徑。
parser.add_argument('--lr', type=float, default=None, help='臨時(shí)修改學(xué)習(xí)率')
parser.add_argument('--epochs', type=int, default=None, help='臨時(shí)修改輪數(shù)') #help是提示信息用于--help的時(shí)候顯示
args = parser.parse_args() #當(dāng)運(yùn)行python train.py --lr 0.001 --epochs 50之后,可以用args.lr調(diào)取這個(gè)值是多少。
# 1. 先加載 yaml 為對象
cfg = load_config_as_obj(args.config) #還是之前的SimpleNamespace。變?yōu)橐粋€(gè)對象,可以用 點(diǎn) 調(diào)用。
# 2. 如果命令行有指定參數(shù),覆蓋 yaml 中的值
if args.lr is not None: #如果通過命令行傳遞進(jìn)來參數(shù)了。
cfg.training.lr = args.lr #重新賦值,進(jìn)而覆蓋Yaml的默認(rèn)值。-這里不會修改yaml文件,只是會修改內(nèi)存里的配置對象cfg
print(f"注意:學(xué)習(xí)率被命令行參數(shù)覆蓋為 {cfg.training.lr}")
if args.epochs is not None:
cfg.training.epochs = args.epochs
return cfg
# 在 main 中調(diào)用:
# cfg = get_args_and_config()
2.json文件
2.1 json文件編寫語法
JSON 是目前互聯(lián)網(wǎng)最通用的數(shù)據(jù)格式。具有語法嚴(yán)格,不能注釋,兼容性較好的特點(diǎn)。
- 語法嚴(yán)格:鍵值對必須用雙引號
""。 - 無注釋:不能寫
#或//,這是它不適合做配置文件的最大原因。 - 兼容性好:網(wǎng)頁、后端、Python 都能直接讀寫。
在深度學(xué)習(xí)中可以存日志 & 存結(jié)果因?yàn)?JSON 機(jī)器讀取速度快且格式標(biāo)準(zhǔn),我們通常用它來保存訓(xùn)練過程中的各項(xiàng)指標(biāo)(Loss, Accuracy),或者數(shù)據(jù)集的標(biāo)注信息(如 COCO 數(shù)據(jù)集)。
json的格式要求更為嚴(yán)格。
- 嚴(yán)謹(jǐn)?shù)逆I值對:類似于 Python 的字典,但要求更嚴(yán)格。
- 雙引號:所有的鍵(Key)和字符串值(Value)必須用雙引號
"",不能用單引號。 - 不支持注釋:這是它最大的特點(diǎn)(也是作為配置文件的缺點(diǎn)),你不能在文件里寫
//或#。 - 數(shù)據(jù)類型:支持 字符串、數(shù)字、布爾值 (
true/false)、列表[]、字典{}。
| 元素 | 寫法要求 | 示例 |
|---|---|---|
| 鍵(key) | 必須是字符串,必須用雙引號包裹 | “batch_size”: 128 |
| 值(value) | 可以是: • 字符串(雙引號) • 數(shù)字 • 布爾值 • null • 對象 • 數(shù)組 | “resnet18” 0.001 true null |
| 字符串 | 必須用雙引號(不能用單引號) | “data_dir”: “./data” |
| 數(shù)字 | 直接寫,不需要引號,支持小數(shù)和科學(xué)計(jì)數(shù)法 | “lr”: 0.001 “weight_decay”: 1e-4 |
| 布爾值 | 只能寫 true 或 false(小寫?。?/td> | “pretrained”: true |
| 空值 | 只能寫 null(小寫) | “optional”: null |
| 數(shù)組 | 用 [ ],元素之間用逗號分隔 | “transforms”: [“RandomCrop”, “Normalize”] |
| 嵌套 | 對象里面可以套對象或數(shù)組 | 見下面的完整例子 |
| 逗號 | 每個(gè)鍵值對或數(shù)組元素后面(除最后一個(gè))必須有逗號 | “batch_size”: 128, |
| 注釋 | 不支持任何形式的注釋(這是和 YAML 最大的區(qū)別?。?/td> | 不能寫 // 或 # 開頭的注釋 |
{
"experiment_id": "exp_2024",
"metrics": {
"accuracy": 0.95,
"loss": 0.045
},
"classes": ["cat", "dog", "car"],
"is_finished": true
}
2.2 json的用法-----類似于yaml
import json
from types import SimpleNamespace # 可選:用來轉(zhuǎn)成點(diǎn)號訪問對象
def load_json_config(path="config.json"):
with open(path, "r", encoding="utf-8") as f:
config_dict = json.load(f) # 注意:是 json.load(f),不是 json.loads()------返回對象也是一個(gè)字典
return config_dict #一個(gè)字典
# 使用
cfg = load_json_config("config.json")
# 字典方式訪問
print(cfg["data"]["batch_size"]) # 128
print(cfg["optimizer"]["lr"]) # 0.001
#轉(zhuǎn)化為對象訪問
def load_json_as_obj(path="config.json"): #給一個(gè)默認(rèn)值config.json,不傳參數(shù)的時(shí)候就用默認(rèn)值。
with open(path, "r", encoding="utf-8") as f:
data = json.load(f)
def dict_to_obj(d):
if isinstance(d, dict): #判斷是不是字典類型
return SimpleNamespace(**{k: dict_to_obj(v) for k, v in d.items()}) #**字典 的意思是:把字典“拆開”成關(guān)鍵字參數(shù)
'''
{ key_expression : value_expression for 變量 in 可迭代對象 } --- {k : dict_to_obj(v) for k, v in d.items()}
d 是一個(gè)字典(比如 {'lr': 0.001, 'device': 'cuda'})
d.items() 返回所有鍵值對:[('lr', 0.001), ('device', 'cuda')]
for k, v in d.items():依次取出鍵(k)和值(v)
k : dict_to_obj(v):新字典的鍵還是原來的 k,但值要先經(jīng)過 dict_to_obj(v) 處理(如果 v 是字典,就遞歸轉(zhuǎn)成對象;如果不是,就原樣返回)
---------------------------
{k: dict_to_obj(v) for k, v in d.items()}-一個(gè)經(jīng)典的字典推導(dǎo)式。它本身就等價(jià)于“先創(chuàng)建一個(gè)空字典,再用 for 循環(huán)往里塞數(shù)據(jù)”。
等價(jià)于:
new_dict = {} # 1 先創(chuàng)建空字典
for k, v in d.items(): # 2 遍歷原字典
new_dict[k] = dict_to_obj(v) # 3 賦值
'''
elif isinstance(d, list): #json可以做嵌套,因此可能包含列表的情況。"classes": ["cat", "dog", "car"],
return [dict_to_obj(i) if isinstance(i, dict) else i for i in d]
'''
new_list = []
for i in d:
if isinstance(i, dict):
new_list.append(dict_to_obj(i))
else:
new_list.append(i)
'''
'''
[表達(dá)式 for 變量 in 可迭代對象 if 條件]
表達(dá)式 → 每次循環(huán)計(jì)算出的值,會成為新列表的元素
變量 → 循環(huán)中取出的每個(gè)元素
可迭代對象 → 任何可遍歷的對象,如列表、字典的 keys、range() 等
if 條件 → 可選,對循環(huán)元素做過濾
'''
else:
return d #普通值
return dict_to_obj(data)
cfg = load_json_as_obj("config.json")
# 現(xiàn)在可以用點(diǎn)號訪問了!
print(cfg.data.batch_size) # 128
print(cfg.optimizer.lr) # 0.001
print(cfg.model.name) # resnet18
#寫入json文件
import json
config = { #創(chuàng)建一個(gè)字典
"project_name": "MyExperiment",
"final_accuracy": 92.5,
"best_epoch": 87
}
with open("result.json", "w", encoding="utf-8") as f: #with open自動打卡文件,用w 模式,with打開不用手動 f.close。with當(dāng)代碼結(jié)束會自動調(diào)用f.close()
json.dump(config, f, indent=4, ensure_ascii=False)
# indent=4:美化輸出,方便閱讀
# ensure_ascii=False:支持中文等非ASCII字符dump函數(shù)講解
json.dump(obj, fp, *, skipkeys=False, ensure_ascii=True, check_circular=True, allow_nan=True, cls=None, indent=None, separators=None, default=None, sort_keys=False)
| 參數(shù) | 類型 | 默認(rèn)值 | 說明 | 推薦用法 |
|---|---|---|---|---|
| obj | Python對象 | 必填 | 要寫入文件的 Python 數(shù)據(jù)(通常是 dict、list、str、int、float、bool、None 等) | 你的配置字典 |
| fp | 文件對象 | 必填 | 已打開的、可寫的文件對象(通常用 open(…, ‘w’)) | with open(…) as f |
| indent | int 或 None | None | 縮進(jìn)空格數(shù)。如果設(shè)置(如 2 或 4),生成的 JSON 會格式化(美化),方便閱讀。每層嵌套增加的空格數(shù),例如每一層嵌套增加 4 個(gè)空格 | indent=4(強(qiáng)烈推薦) |
| ensure_ascii | bool | True | 如果為 True,非 ASCII 字符(如中文)會轉(zhuǎn)成 \uXXXX 轉(zhuǎn)義。如果為 False,直接保留原字符 | ensure_ascii=False(有中文時(shí)必設(shè)) |
| sort_keys | bool | False | 是否對字典的鍵進(jìn)行排序(按字母順序) | sort_keys=True(調(diào)試時(shí)方便對比) |
| separators | tuple | (', ', ': ') | 控制項(xiàng)分隔符和鍵值分隔符,通常不用改 | 一般不改 |
| default | callable | None | 如果對象有無法序列化的類型(如 set、datetime),可以用這個(gè)函數(shù)自定義轉(zhuǎn)換 | 高級用法 |
3.py文件—實(shí)現(xiàn)“代碼即配置”
把配置從“純數(shù)據(jù)”(data)升級成“可執(zhí)行代碼”(code)。 簡單說,就是直接用一個(gè) Python 文件(通常叫 config.py、models_config.py 等)來定義所有配置和邏輯,而不是用 YAML/JSON 只存靜態(tài)值。這在深度學(xué)習(xí)項(xiàng)目中非常常見,尤其是當(dāng)配置需要包含復(fù)雜邏輯時(shí)(比如動態(tài)構(gòu)建模型、條件判斷、計(jì)算路徑等)。
# config.py - 所有配置和邏輯集中在這里
import torch
import torch.nn as nn
from torchvision import models, transforms
# ================== 數(shù)據(jù)配置 ==================
DATA_DIR = "./data"
BATCH_SIZE = 128
NUM_WORKERS = 4
# ================== 模型配置 ==================
MODEL_NAME = "resnet18" # 改這里就能換模型
NUM_CLASSES = 10
PRETRAINED = True
def get_model():
if MODEL_NAME == "resnet18":
base = models.resnet18(pretrained=PRETRAINED)
elif MODEL_NAME == "resnet50":
base = models.resnet50(pretrained=PRETRAINED)
elif MODEL_NAME == "mobilenet_v2":
base = models.mobilenet_v2(pretrained=PRETRAINED)
else:
raise ValueError(f"Unknown model: {MODEL_NAME}")
# 統(tǒng)一修改最后一層
if hasattr(base, 'fc'): # ResNet 系列
base.fc = nn.Linear(base.fc.in_features, NUM_CLASSES)
elif hasattr(base, 'classifier'): # MobileNet
base.classifier[1] = nn.Linear(base.classifier[1].in_features, NUM_CLASSES)
return base
# ================== 訓(xùn)練配置 ==================
LR = 0.001
EPOCHS = 100
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
# ================== 數(shù)據(jù)增強(qiáng) ==================
def get_transforms():
return transforms.Compose([
transforms.RandomCrop(32, padding=4),
transforms.RandomHorizontalFlip(),
transforms.ToTensor(),
transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010)),
])
# train.py from config import get_model, get_transforms, BATCH_SIZE, NUM_WORKERS, LR, EPOCHS, DEVICE, DATA_DIR model = get_model().to(DEVICE) transform = get_transforms() # 數(shù)據(jù)加載、優(yōu)化器、訓(xùn)練循環(huán)...
4.XML
XML(eXtensible Markup Language) 是一種 可擴(kuò)展標(biāo)記語言,用于存儲和傳輸數(shù)據(jù),類似 JSON/YAML。具有一下特點(diǎn):
- 可讀性強(qiáng),層級清晰
- 支持嵌套和屬性
但是比 JSON/YAML 冗長,而且使用起來有點(diǎn)復(fù)雜。
在目標(biāo)檢測(Object Detection)領(lǐng)域,尤其是經(jīng)典的 Pascal VOC 數(shù)據(jù)集(2007/2012)和很多自定義數(shù)據(jù)集,標(biāo)注信息都是用 XML 文件 來存儲的。 每一個(gè)圖像對應(yīng)一個(gè) .xml 文件,里面記錄了圖像中所有目標(biāo)的類別、邊界框坐標(biāo)(bounding box)等信息。
<?xml version="1.0" encoding="UTF-8"?>
<config>
<project_name>MyExperiment</project_name>
<model>
<type>resnet18</type>
<num_classes>10</num_classes>
<pretrained>true</pretrained>
</model>
<training>
<batch_size>64</batch_size>
<epochs>100</epochs>
<learning_rate>0.001</learning_rate>
<optimizer>Adam</optimizer>
</training>
</config>
<config>→ 根節(jié)點(diǎn)(root element)<model>/<training>→ 子節(jié)點(diǎn)<type>resnet18</type>→ 標(biāo)簽 + 內(nèi)容- XML 支持 嵌套層級,適合復(fù)雜配置
4.1讀取
#Python 內(nèi)置庫 xml.etree.ElementTree 可以解析 XML。解析、創(chuàng)建、操作 XML 文件
import xml.etree.ElementTree as ET
# 1. 讀取 XML 文件
tree = ET.parse("config.xml") # 返回 ElementTree 對象 --tree → XML 的整個(gè)樹形結(jié)構(gòu)
root = tree.getroot() # 根節(jié)點(diǎn) <config>----從 ElementTree 中獲取 根節(jié)點(diǎn)
# 2. 訪問數(shù)據(jù)
project_name = root.find("project_name").text #root.find--查找 <config> 下的第一個(gè) <project_name> 子節(jié)點(diǎn),返回一個(gè) Element 對象
#.text 獲取該節(jié)點(diǎn)的文本內(nèi)容 "MyExperiment"。
print(project_name) # MyExperiment---project_name → 字符串類
model_type = root.find("model/type").text
num_classes = int(root.find("model/num_classes").text)
pretrained = root.find("model/pretrained").text == "true" #== "true" → 轉(zhuǎn)成布爾值
batch_size = int(root.find("training/batch_size").text)
learning_rate = float(root.find("training/learning_rate").text)
print(model_type, num_classes, pretrained, batch_size, learning_rate)
5.TOML
TOML(Tom’s Obvious, Minimal Language)是一種現(xiàn)代、簡潔、人性化的配置文件格式,由 GitHub 聯(lián)合創(chuàng)始人 Tom Preston-Werner 創(chuàng)建。它的設(shè)計(jì)目標(biāo)是盡可能明顯、直觀,比 JSON 可讀性更強(qiáng)(支持注釋),比 YAML 更簡單(縮進(jìn)不敏感)。在深度學(xué)習(xí)項(xiàng)目中,TOML 的最主流、最核心用法不是在代碼里讀寫超參數(shù),而是用于項(xiàng)目依賴管理和構(gòu)建配置——即 pyproject.toml 文件。
5.1 TOML的語法格式
| 元素類型 | 寫法要求 | 示例 | 說明 |
|---|---|---|---|
| 鍵值對 | key = value(等號兩邊有空格) | batch_size = 128 | 最基本的配置方式 |
| 字符串 | 單引號 '...' 或雙引號 "..." | name = "resnet18" | 可使用轉(zhuǎn)義字符 |
| 數(shù)字 | 直接寫,支持整數(shù)、浮點(diǎn)數(shù)、科學(xué)計(jì)數(shù)法 | lr = 0.001、weight_decay = 1e-4 | 默認(rèn)是數(shù)字類型 |
| 布爾值 | true 或 false(小寫) | pretrained = true | XML/JSON 沒有布爾類型要特別注意 |
| 數(shù)組/列表 | [elem1, elem2, ...] | transforms = ["crop", "flip"] | 支持不同類型混合元素 |
| 表(Table) | [table_name] 或點(diǎn)號嵌套 table.subtable | [model] 或 model.name = "resnet18" | 用于分組或嵌套配置 |
| 注釋 | # 開頭 | # 這是注釋 | 注釋不會被解析 |
| 多行字符串 | 三個(gè)引號 """...""" | desc = """多行文本""" | 支持換行 |
| 嵌套表 | [table.subtable] 或 table.subtable.key = value | [data.train] | 支持多層嵌套結(jié)構(gòu) |
| 日期/時(shí)間 | ISO 8601 格式 | start_date = 2025-12-24T22:00:00Z | TOML 內(nèi)置日期時(shí)間類型 |
# config.toml project_name = "CIFAR10_Classification" seed = 42 [data] dataset = "CIFAR10" data_dir = "./data" batch_size = 128 num_workers = 4 [model] #TOML 使用 表(table) 來表示 嵌套結(jié)構(gòu)或命名空間,[model] 表示 一個(gè)名為 model 的表,表下面的鍵值對都屬于這個(gè)表的 子空間 name = "resnet18" pretrained = true num_classes = 10 [optimizer] name = "Adam" lr = 0.001 weight_decay = 1e-4 [train] epochs = 100 device = "cuda"
5.2 toml的用法
#可以做項(xiàng)目依賴管理
'''
TOML 的 最常見用途不是存超參,而是 管理 Python 項(xiàng)目的依賴和構(gòu)建配置
在現(xiàn)代 Python 項(xiàng)目中,它已經(jīng)取代了:
requirements.txt(老式依賴列表)
setup.py(舊版打包配置)
存放位置:項(xiàng)目根目錄
'''
#例如
[project]
name = "cifar10-resnet"
version = "0.1.0"
description = "CIFAR-10 分類實(shí)驗(yàn)"
authors = [{name = "張三", email = "zhangsan@example.com"}]
requires-python = ">=3.9"
dependencies = [
"torch>=2.0.0",
"torchvision>=0.15.0",
"pyyaml>=6.0",
"matplotlib>=3.5",
"tqdm"
]
[project.optional-dependencies]
dev = ["black", "flake8", "pytest"]
train = ["wandb", "tensorboard"]
[build-system]
requires = ["hatchling"]
build-backend = "hatchling.build"
#可以這么用
# 安裝項(xiàng)目依賴
pip install .
# 安裝開發(fā)依賴
pip install -e ".[dev]"
# 用 poetry 管理
poetry install # 自動安裝所有依賴
poetry add torch==2.1.0 # 添加新依賴,自動更新 toml
'''
環(huán)境可復(fù)現(xiàn):別人 clone 你的代碼后,只需 pip install . 就能裝好相同版本的包
版本鎖定:精確控制 torch、torchvision 等版本
現(xiàn)代標(biāo)準(zhǔn):pip、poetry、pdm 等工具都支持 TOML
分組依賴:區(qū)分運(yùn)行、開發(fā)、訓(xùn)練依賴
在工程中,TOML 最重要的用途是 依賴管理和項(xiàng)目構(gòu)建,而不是超參配置
'''#作為超參配置--代碼讀取 TOML 作為超參配置
#雖然不常用,但可以把 TOML 當(dāng)作 YAML/JSON 的替代品,存超參。需要安裝toml庫。pip install toml
#讀取toml為對象
import toml
from types import SimpleNamespace
def load_toml_config(path="config.toml"):
data = toml.load(path) #返回字典類型
def dict_to_obj(d):
if isinstance(d, dict):
return SimpleNamespace(**{k: dict_to_obj(v) for k, v in d.items()})
elif isinstance(d, list): #可能有淚飆類型 transforms = ["crop", "flip", "normalize"] 讓列表里的字典也能用 點(diǎn)號訪問。
return [dict_to_obj(i) for i in d]
else:
return d
return dict_to_obj(data)
cfg = load_toml_config("config.toml")
print(cfg.data.batch_size) # 128
print(cfg.optimizer.lr) # 0.001
#寫入toml
import toml
config = {"train": {"epochs": 100, "lr": 0.001}}
with open("config.toml", "w") as f:
toml.dump(config, f)
以上就是Pytorch中項(xiàng)目配置文件的管理與導(dǎo)入方式的詳細(xì)內(nèi)容,更多關(guān)于Pytorch配置文件管理與導(dǎo)入的資料請關(guān)注腳本之家其它相關(guān)文章!
相關(guān)文章
安裝出現(xiàn):Requirement?already?satisfied解決辦法
最近pip install的時(shí)候報(bào)錯(cuò),一大串Requirement already satisfied,所以下面這篇文章主要給大家介紹了關(guān)于安裝出現(xiàn):Requirement?already?satisfied的解決辦法,需要的朋友可以參考下2022-08-08
Python實(shí)現(xiàn)登陸文件驗(yàn)證方法
本篇文章中我們給大家分享了關(guān)于Python實(shí)現(xiàn)登陸文件驗(yàn)證的方法和技巧,有興趣的朋友們參考學(xué)習(xí)下。2018-10-10
關(guān)于使用OpenCsv導(dǎo)入大數(shù)據(jù)量報(bào)錯(cuò)的問題
這篇文章主要介紹了使用OpenCsv導(dǎo)入大數(shù)據(jù)量報(bào)錯(cuò)的問題 ,本文給大家介紹的非常詳細(xì),對大家的學(xué)習(xí)或工作具有一定的參考借鑒價(jià)值,需要的朋友可以參考下2021-08-08
深入解析Python中delattr函數(shù)的使用方法和應(yīng)用場景
這篇文章主要為大家詳細(xì)介紹了Python中delattr()函數(shù)的使用方法和應(yīng)用場景,文中的示例代碼講解詳細(xì),感興趣的小伙伴可以跟隨小編一起學(xué)習(xí)一下2026-02-02
在PyCharm導(dǎo)航區(qū)中打開多個(gè)Project的關(guān)閉方法
今天小編就為大家分享一篇在PyCharm導(dǎo)航區(qū)中打開多個(gè)Project的關(guān)閉方法,具有很好的參考價(jià)值,希望對大家有所幫助。一起跟隨小編過來看看吧2019-01-01
在pycharm 中添加運(yùn)行參數(shù)的操作方法
今天小編就為大家分享一篇在pycharm 中添加運(yùn)行參數(shù)的操作方法,具有很好的參考價(jià)值,希望對大家有所幫助。一起跟隨小編過來看看吧2019-01-01

