| 1217 | } |
| 1218 | |
| 1219 | TF_AttrMetadata TF_OperationGetAttrMetadata(TF_Operation* oper, |
| 1220 | const char* attr_name, |
| 1221 | TF_Status* status) { |
| 1222 | TF_AttrMetadata metadata; |
| 1223 | const auto* attr = GetAttrValue(oper, attr_name, status); |
| 1224 | if (TF_GetCode(status) != TF_OK) return metadata; |
| 1225 | switch (attr->value_case()) { |
| 1226 | #define SINGLE_CASE(kK, attr_type, size_expr) \ |
| 1227 | case tensorflow::AttrValue::kK: \ |
| 1228 | metadata.is_list = 0; \ |
| 1229 | metadata.list_size = -1; \ |
| 1230 | metadata.type = attr_type; \ |
| 1231 | metadata.total_size = size_expr; \ |
| 1232 | break; |
| 1233 | |
| 1234 | SINGLE_CASE(kS, TF_ATTR_STRING, attr->s().length()); |
| 1235 | SINGLE_CASE(kI, TF_ATTR_INT, -1); |
| 1236 | SINGLE_CASE(kF, TF_ATTR_FLOAT, -1); |
| 1237 | SINGLE_CASE(kB, TF_ATTR_BOOL, -1); |
| 1238 | SINGLE_CASE(kType, TF_ATTR_TYPE, -1); |
| 1239 | SINGLE_CASE(kShape, TF_ATTR_SHAPE, |
| 1240 | attr->shape().unknown_rank() ? -1 : attr->shape().dim_size()); |
| 1241 | SINGLE_CASE(kTensor, TF_ATTR_TENSOR, -1); |
| 1242 | #undef SINGLE_CASE |
| 1243 | |
| 1244 | case tensorflow::AttrValue::kList: |
| 1245 | metadata.is_list = 1; |
| 1246 | metadata.list_size = 0; |
| 1247 | metadata.total_size = -1; |
| 1248 | #define LIST_CASE(field, attr_type, ...) \ |
| 1249 | if (attr->list().field##_size() > 0) { \ |
| 1250 | metadata.type = attr_type; \ |
| 1251 | metadata.list_size = attr->list().field##_size(); \ |
| 1252 | __VA_ARGS__; \ |
| 1253 | break; \ |
| 1254 | } |
| 1255 | |
| 1256 | LIST_CASE( |
| 1257 | s, TF_ATTR_STRING, metadata.total_size = 0; |
| 1258 | for (int i = 0; i < attr->list().s_size(); |
| 1259 | ++i) { metadata.total_size += attr->list().s(i).size(); }); |
| 1260 | LIST_CASE(i, TF_ATTR_INT); |
| 1261 | LIST_CASE(f, TF_ATTR_FLOAT); |
| 1262 | LIST_CASE(b, TF_ATTR_BOOL); |
| 1263 | LIST_CASE(type, TF_ATTR_TYPE); |
| 1264 | LIST_CASE( |
| 1265 | shape, TF_ATTR_SHAPE, metadata.total_size = 0; |
| 1266 | for (int i = 0; i < attr->list().shape_size(); ++i) { |
| 1267 | const auto& s = attr->list().shape(i); |
| 1268 | metadata.total_size += s.unknown_rank() ? 0 : s.dim_size(); |
| 1269 | }); |
| 1270 | LIST_CASE(tensor, TF_ATTR_TENSOR); |
| 1271 | LIST_CASE(tensor, TF_ATTR_FUNC); |
| 1272 | #undef LIST_CASE |
| 1273 | // All lists empty, determine the type from the OpDef. |
| 1274 | if (metadata.list_size == 0) { |
| 1275 | for (int i = 0; i < oper->node.op_def().attr_size(); ++i) { |
| 1276 | const auto& a = oper->node.op_def().attr(i); |