-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmoco_parameter_aliases.py
More file actions
245 lines (191 loc) · 6.2 KB
/
Copy pathmoco_parameter_aliases.py
File metadata and controls
245 lines (191 loc) · 6.2 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
"""
MoCo参数别名映射配置
本模块定义了MoCo参数的别名映射关系,用于处理参数统一过程中的向后兼容性。
"""
from typing import Dict, Any
# MoCo参数别名映射表
# 格式: {别名: 标准参数名}
MOCO_PARAMETER_ALIASES = {
# 动量系数别名
'moco_m': 'moco_momentum',
# 温度系数别名
'moco_T': 'moco_t',
# 可能的未来别名扩展
'momentum': 'moco_momentum',
'temperature': 'moco_t',
'tau': 'moco_t',
'queue_size': 'moco_K',
'queue_length': 'moco_queue'
}
# 反向映射表(标准参数名到所有可能的别名)
MOCO_PARAMETER_REVERSE_ALIASES = {}
for alias, standard in MOCO_PARAMETER_ALIASES.items():
if standard not in MOCO_PARAMETER_REVERSE_ALIASES:
MOCO_PARAMETER_REVERSE_ALIASES[standard] = []
MOCO_PARAMETER_REVERSE_ALIASES[standard].append(alias)
# MoCo参数默认值配置
MOCO_PARAMETER_DEFAULTS = {
'moco_momentum': 0.999,
'moco_t': 0.2,
'moco_tau1': 0.2,
'moco_tau2': 0.3,
'moco_K': 4096,
'moco_queue': 4096,
'enable_view_0': True,
'moco_type': 'basic',
'proj_dim': None,
'queue_warmup_steps': 0
}
# MoCo参数类型定义
MOCO_PARAMETER_TYPES = {
'moco_momentum': float,
'moco_t': float,
'moco_tau1': float,
'moco_tau2': float,
'moco_K': int,
'moco_queue': int,
'enable_view_0': bool,
'moco_type': str,
'proj_dim': int,
'queue_warmup_steps': int
}
# MoCo参数范围定义
MOCO_PARAMETER_RANGES = {
'moco_momentum': (0.9, 0.9999),
'moco_t': (0.01, 1.0),
'moco_tau1': (0.01, 1.0),
'moco_tau2': (0.01, 1.0),
'moco_K': [1024, 2048, 4096, 8192],
'moco_queue': [1024, 2048, 4096, 8192],
'enable_view_0': ['true', 'false'],
'moco_type': ['basic', 'double_tau'],
'proj_dim': (32, 512),
'queue_warmup_steps': (0, 1000)
}
def resolve_parameter_alias(param_name: str) -> str:
"""
解析参数别名,返回标准参数名
Args:
param_name: 参数名(可能是别名)
Returns:
标准参数名
"""
return MOCO_PARAMETER_ALIASES.get(param_name, param_name)
def apply_parameter_aliases(parameters: Dict[str, Any]) -> Dict[str, Any]:
"""
应用参数别名映射,将别名转换为标准参数名
Args:
parameters: 原始参数字典
Returns:
转换后的参数字典
"""
resolved_params = {}
for param_name, value in parameters.items():
standard_name = resolve_parameter_alias(param_name)
resolved_params[standard_name] = value
return resolved_params
def get_parameter_default(param_name: str) -> Any:
"""
获取参数的默认值
Args:
param_name: 参数名
Returns:
参数默认值,如果参数不存在则返回None
"""
standard_name = resolve_parameter_alias(param_name)
return MOCO_PARAMETER_DEFAULTS.get(standard_name)
def get_parameter_type(param_name: str) -> type:
"""
获取参数的类型
Args:
param_name: 参数名
Returns:
参数类型,如果参数不存在则返回None
"""
standard_name = resolve_parameter_alias(param_name)
return MOCO_PARAMETER_TYPES.get(standard_name)
def get_parameter_range(param_name: str) -> Any:
"""
获取参数的取值范围
Args:
param_name: 参数名
Returns:
参数取值范围(元组表示连续范围,列表表示离散值)
"""
standard_name = resolve_parameter_alias(param_name)
return MOCO_PARAMETER_RANGES.get(standard_name)
def validate_moco_parameter(param_name: str, value: Any) -> bool:
"""
验证MoCo参数值是否有效
Args:
param_name: 参数名
value: 参数值
Returns:
True如果参数值有效,否则False
"""
standard_name = resolve_parameter_alias(param_name)
param_range = get_parameter_range(standard_name)
param_type = get_parameter_type(standard_name)
if param_range is None or param_type is None:
return False
try:
# 类型转换验证
if param_type == bool:
if isinstance(value, str):
value = value.lower() in ['true', '1', 'yes', 'on']
else:
value = bool(value)
else:
value = param_type(value)
# 范围验证
if isinstance(param_range, tuple):
# 连续范围
return param_range[0] <= value <= param_range[1]
elif isinstance(param_range, list):
# 离散值
return value in param_range
except (ValueError, TypeError):
return False
return False
def get_all_moco_parameters() -> list:
"""
获取所有MoCo参数的标准名称列表
Returns:
MoCo参数名称列表
"""
return list(MOCO_PARAMETER_DEFAULTS.keys())
def is_moco_parameter(param_name: str) -> bool:
"""
检查参数是否为MoCo相关参数
Args:
param_name: 参数名
Returns:
True如果是MoCo参数,否则False
"""
standard_name = resolve_parameter_alias(param_name)
return standard_name in MOCO_PARAMETER_DEFAULTS
if __name__ == "__main__":
# 测试代码
print("测试MoCo参数别名映射...")
# 测试别名解析
test_params = {
'moco_m': 0.999,
'moco_T': 0.2,
'moco_tau1': 0.15,
'enable_view_0': 'true'
}
print(f"原始参数: {test_params}")
resolved = apply_parameter_aliases(test_params)
print(f"解析后参数: {resolved}")
# 测试参数验证
for param, value in resolved.items():
is_valid = validate_moco_parameter(param, value)
print(f"参数 {param}={value} 验证结果: {is_valid}")
# 测试默认值获取
print(f"\n默认值测试:")
for param in get_all_moco_parameters():
default = get_parameter_default(param)
param_type = get_parameter_type(param)
param_range = get_parameter_range(param)
print(f" {param}: 默认={default}, 类型={param_type.__name__}, 范围={param_range}")
print("\nMoCo参数别名映射测试完成!")