compute_gradiend_feature_cross_encoding_matrix
compute_gradiend_feature_cross_encoding_matrix
compute_gradiend_feature_cross_encoding_matrix(trainers, feature_classes, *, trainer_order=None, eval_by_class=None, split='test', max_size=None)
Compute a dense GRADIEND by feature-class cross-encoding matrix.
Row i is a trained GRADIEND; column j is a feature class. Cell (i, j)
is the mean encoded value when GRADIEND i encodes shared eval snippets for
class j. When eval_by_class is omitted, snippets are collected from the
trainers' unified data for the requested split.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
trainers
|
Dict[str, object]
|
Mapping from trainer id to trainer object. Trainers must be able to load a model and create a gradient training dataset for the generated probe pairs. |
required |
feature_classes
|
Sequence[str]
|
Column order for feature classes to evaluate. |
required |
trainer_order
|
Optional[Sequence[str]]
|
Optional row order. Unknown ids are ignored; at least one valid id must remain. |
None
|
eval_by_class
|
Optional[Dict[str, DataFrame]]
|
Optional precomputed mapping from feature class to a
unified-data DataFrame. If omitted, it is built from |
None
|
split
|
str
|
Split used when collecting unified eval rows. |
'test'
|
max_size
|
Optional[int]
|
Optional maximum examples per feature class. If set, rows are sampled with a fixed random seed before encoding. |
None
|
Returns:
| Type | Description |
|---|---|
Dict[str, Any]
|
A payload with |
Dict[str, Any]
|
|
Dict[str, Any]
|
|
Dict[str, Any]
|
|
Raises:
| Type | Description |
|---|---|
ValueError
|
If no trainers/classes are provided, |
Source code in gradiend/comparison/cross_encoding.py
1231 1232 1233 1234 1235 1236 1237 1238 1239 1240 1241 1242 1243 1244 1245 1246 1247 1248 1249 1250 1251 1252 1253 1254 1255 1256 1257 1258 1259 1260 1261 1262 1263 1264 1265 1266 1267 1268 1269 1270 1271 1272 1273 1274 1275 1276 1277 1278 1279 1280 1281 1282 1283 1284 1285 1286 1287 1288 1289 1290 1291 1292 1293 1294 1295 1296 1297 1298 1299 1300 1301 1302 1303 1304 1305 1306 1307 1308 1309 1310 1311 1312 1313 1314 1315 1316 1317 1318 1319 1320 1321 1322 1323 1324 | |