| 399 | } |
| 400 | |
| 401 | string XlaCompiler::Argument::HumanString() const { |
| 402 | string common; |
| 403 | if (!name.empty()) { |
| 404 | common = absl::StrCat(" name=", name); |
| 405 | } |
| 406 | absl::StrAppend(&common, " type=", DataTypeString(type), |
| 407 | " shape=", ShapeHumanString()); |
| 408 | absl::StrAppend( |
| 409 | &common, " is_same_data_across_replicas=", is_same_data_across_replicas); |
| 410 | switch (kind) { |
| 411 | case kInvalid: |
| 412 | return "invalid"; |
| 413 | case kConstant: |
| 414 | return absl::StrCat("kind=constant", common, |
| 415 | " value=", constant_value.DebugString()); |
| 416 | case kResource: { |
| 417 | string output = absl::StrCat("kind=resource", common, " resource_kind=", |
| 418 | XlaResource::KindToString(resource_kind), |
| 419 | " initialized=", initialized); |
| 420 | if (max_array_size >= 0) { |
| 421 | absl::StrAppend(&output, " max_array_size=", max_array_size); |
| 422 | } |
| 423 | if (!tensor_array_gradients.empty()) { |
| 424 | absl::StrAppend(&output, " tensor_array_gradients=", |
| 425 | absl::StrJoin(tensor_array_gradients, ",")); |
| 426 | } |
| 427 | return output; |
| 428 | } |
| 429 | case kParameter: |
| 430 | return absl::StrCat("kind=parameter", common); |
| 431 | case kTensorList: |
| 432 | return absl::StrCat("kind=tensorlist", common); |
| 433 | case kToken: |
| 434 | return absl::StrCat("token", common); |
| 435 | } |
| 436 | } |
| 437 | |
| 438 | std::vector<int64> XlaCompiler::Argument::DimensionSizes() const { |
| 439 | if (absl::holds_alternative<TensorShape>(shape)) { |