Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 11 additions & 0 deletions tensorflow/compiler/jit/encapsulate_subgraphs_pass.cc
Original file line number Diff line number Diff line change
Expand Up @@ -139,6 +139,17 @@ std::string ExprProtoToString(const ExpressionProto& e) {
case ExpressionProto::kDivNode:
return absl::StrCat("(", ExprProtoToString(e.div_node().lhs()), " / ",
ExprProtoToString(e.div_node().rhs()), ")");
case ExpressionProto::kMaxNode:
return absl::StrCat("max(", ExprProtoToString(e.max_node().lhs()), ", ",
ExprProtoToString(e.max_node().rhs()), ")");
case ExpressionProto::kGtNode:
return absl::StrCat("(", ExprProtoToString(e.gt_node().lhs()), " > ",
ExprProtoToString(e.gt_node().rhs()), ")");
case ExpressionProto::kSelectNode:
return absl::StrCat("select(", ExprProtoToString(e.select_node().pred()),
", ", ExprProtoToString(e.select_node().on_true()),
", ", ExprProtoToString(e.select_node().on_false()),
")");
default:
return "<none>";
}
Expand Down
80 changes: 54 additions & 26 deletions tensorflow/compiler/jit/kernels/xla_ops.cc
Original file line number Diff line number Diff line change
Expand Up @@ -101,6 +101,8 @@ limitations under the License.
namespace tensorflow {

namespace {
constexpr char kUserInferredValueContentsAttrName[] =
"_user_inferred_value_contents";
using XlaDeviceCompiler =
DeviceCompiler<xla::LocalExecutable, xla::LocalClient>;
using PjRtDeviceCompiler =
Expand Down Expand Up @@ -414,6 +416,23 @@ std::unique_ptr<DimExpr> ExprFromProto(const ExpressionProto& proto) {
auto rhs = ExprFromProto(proto.div_node().rhs());
return std::make_unique<ExprDiv>(lhs.release(), rhs.release());
}
case ExpressionProto::kMaxNode: {
auto lhs = ExprFromProto(proto.max_node().lhs());
auto rhs = ExprFromProto(proto.max_node().rhs());
return std::make_unique<ExprMax>(lhs.release(), rhs.release());
}
case ExpressionProto::kGtNode: {
auto lhs = ExprFromProto(proto.gt_node().lhs());
auto rhs = ExprFromProto(proto.gt_node().rhs());
return std::make_unique<ExprGt>(lhs.release(), rhs.release());
}
case ExpressionProto::kSelectNode: {
auto pred = ExprFromProto(proto.select_node().pred());
auto on_true = ExprFromProto(proto.select_node().on_true());
auto on_false = ExprFromProto(proto.select_node().on_false());
return std::make_unique<ExprSelect>(pred.release(), on_true.release(),
on_false.release());
}
case ExpressionProto::NODE_TYPE_NOT_SET:
default:
return nullptr;
Expand Down Expand Up @@ -445,6 +464,22 @@ static xla::DExpr DimExprToDExpr(const DimExpr* e) {
auto* ee = static_cast<const ExprDiv*>(e);
return DimExprToDExpr(ee->lhs()) / DimExprToDExpr(ee->rhs());
}
case DimExpr::Kind::kMax: {
auto* ee = static_cast<const ExprMax*>(e);
return xla::DExpr::Max(DimExprToDExpr(ee->lhs()),
DimExprToDExpr(ee->rhs()));
}
case DimExpr::Kind::kGt: {
auto* ee = static_cast<const ExprGt*>(e);
return xla::DExpr::Gt(DimExprToDExpr(ee->lhs()),
DimExprToDExpr(ee->rhs()));
}
case DimExpr::Kind::kSelect: {
auto* ee = static_cast<const ExprSelect*>(e);
return xla::DExpr::Select(DimExprToDExpr(ee->pred()),
DimExprToDExpr(ee->on_true()),
DimExprToDExpr(ee->on_false()));
}
}
return xla::DExpr::Unknown();
}
Expand Down Expand Up @@ -520,35 +555,35 @@ absl::Status CompileToLocalExecutable(
return;
}

auto inferred_shape_it = attr_map.find("user_inferred_shape");
auto inferred_contents_it =
attr_map.find(kUserInferredValueContentsAttrName);
bool has_dynamic = false;
auto has_dynamic_it = attr_map.find("has_dynamic");
if (has_dynamic_it == attr_map.end()) {
return;
}
has_dynamic = has_dynamic_it->second.b();
if (!has_dynamic) {
return;
if (has_dynamic_it != attr_map.end()) {
has_dynamic = has_dynamic_it->second.b();
}

auto inferred_shape_it = attr_map.find("user_inferred_shape");
if (inferred_shape_it == attr_map.end()) {
VLOG(1) << "XlaCompileOp saw has_dynamic for const arg "
<< arg_index << " node=" << node_name
<< " but no user_inferred_shape attr";
if (inferred_contents_it == attr_map.end() &&
(!has_dynamic || inferred_shape_it == attr_map.end())) {
return;
}

TensorShapeProto inferred_shape_proto;
inferred_shape_proto = inferred_shape_it->second.shape();
if (inferred_contents_it != attr_map.end()) {
if (!inferred_shape_proto.ParseFromString(
inferred_contents_it->second.s())) {
return;
}
} else {
inferred_shape_proto = inferred_shape_it->second.shape();
}

TensorShape inferred_shape(inferred_shape_proto);
if (!TensorShapeUtils::IsVector(arg.constant_value.shape()) ||
arg.constant_value.NumElements() != inferred_shape.dims()) {
VLOG(1) << "XlaCompileOp const arg " << arg_index
<< " node=" << node_name
<< " has dynamic shape metadata but tensor shape "
<< arg.constant_value.shape().DebugString()
<< " does not match inferred rank " << inferred_shape.dims();
if (!((TensorShapeUtils::IsVector(arg.constant_value.shape()) &&
arg.constant_value.NumElements() == inferred_shape.dims()) ||
(TensorShapeUtils::IsScalar(arg.constant_value.shape()) &&
inferred_shape.dims() == 1))) {
return;
}

Expand All @@ -564,18 +599,11 @@ absl::Status CompileToLocalExecutable(
} else if (arg.constant_value.dtype() == DT_INT64) {
expr.set_constant_value(arg.constant_value.flat<int64_t>()(i));
} else {
VLOG(1) << "XlaCompileOp const arg " << arg_index
<< " node=" << node_name
<< " has unsupported dtype for inferred shape contents: "
<< DataTypeString(arg.constant_value.dtype());
arg.constant_value_expressions.clear();
return;
}
arg.constant_value_expressions.push_back(std::move(expr));
}
VLOG(1) << "XlaCompileOp recovered " << arg.constant_value_expressions.size()
<< " constant_value_expressions for const arg " << arg_index
<< " node=" << node_name << " from user_inferred_shape";
};
auto record_dynamic_dim_value = [&](int64_t dim_size, xla::DExpr expr) {
if (!saw_dynamic_dim_value) {
Expand Down
44 changes: 44 additions & 0 deletions tensorflow/compiler/jit/mark_for_compilation_pass.cc
Original file line number Diff line number Diff line change
Expand Up @@ -714,6 +714,17 @@ std::string ExprProtoToString(const ExpressionProto& e) {
case ExpressionProto::kDivNode:
return absl::StrCat("(", ExprProtoToString(e.div_node().lhs()), " / ",
ExprProtoToString(e.div_node().rhs()), ")");
case ExpressionProto::kMaxNode:
return absl::StrCat("max(", ExprProtoToString(e.max_node().lhs()), ", ",
ExprProtoToString(e.max_node().rhs()), ")");
case ExpressionProto::kGtNode:
return absl::StrCat("(", ExprProtoToString(e.gt_node().lhs()), " > ",
ExprProtoToString(e.gt_node().rhs()), ")");
case ExpressionProto::kSelectNode:
return absl::StrCat("select(", ExprProtoToString(e.select_node().pred()),
", ", ExprProtoToString(e.select_node().on_true()),
", ", ExprProtoToString(e.select_node().on_false()),
")");
default:
return "<none>";
}
Expand Down Expand Up @@ -747,6 +758,23 @@ std::unique_ptr<DimExpr> ExprFromProto(const ExpressionProto& proto) {
auto rhs = ExprFromProto(proto.div_node().rhs());
return std::make_unique<ExprDiv>(lhs.release(), rhs.release());
}
case ExpressionProto::kMaxNode: {
auto lhs = ExprFromProto(proto.max_node().lhs());
auto rhs = ExprFromProto(proto.max_node().rhs());
return std::make_unique<ExprMax>(lhs.release(), rhs.release());
}
case ExpressionProto::kGtNode: {
auto lhs = ExprFromProto(proto.gt_node().lhs());
auto rhs = ExprFromProto(proto.gt_node().rhs());
return std::make_unique<ExprGt>(lhs.release(), rhs.release());
}
case ExpressionProto::kSelectNode: {
auto pred = ExprFromProto(proto.select_node().pred());
auto on_true = ExprFromProto(proto.select_node().on_true());
auto on_false = ExprFromProto(proto.select_node().on_false());
return std::make_unique<ExprSelect>(pred.release(), on_true.release(),
on_false.release());
}
case ExpressionProto::NODE_TYPE_NOT_SET:
default:
return nullptr;
Expand Down Expand Up @@ -779,6 +807,22 @@ static xla::DExpr DimExprToDExpr(const DimExpr* e) {
auto* ee = static_cast<const ExprDiv*>(e);
return DimExprToDExpr(ee->lhs()) / DimExprToDExpr(ee->rhs());
}
case DimExpr::Kind::kMax: {
auto* ee = static_cast<const ExprMax*>(e);
return xla::DExpr::Max(DimExprToDExpr(ee->lhs()),
DimExprToDExpr(ee->rhs()));
}
case DimExpr::Kind::kGt: {
auto* ee = static_cast<const ExprGt*>(e);
return xla::DExpr::Gt(DimExprToDExpr(ee->lhs()),
DimExprToDExpr(ee->rhs()));
}
case DimExpr::Kind::kSelect: {
auto* ee = static_cast<const ExprSelect*>(e);
return xla::DExpr::Select(DimExprToDExpr(ee->pred()),
DimExprToDExpr(ee->on_true()),
DimExprToDExpr(ee->on_false()));
}
}
return xla::DExpr();
}
Expand Down
32 changes: 28 additions & 4 deletions tensorflow/compiler/tf2xla/kernels/batchtospace_op.cc
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,8 @@ void BatchToSpace(XlaOpKernelContext* ctx, const xla::XlaOp input,
const int input_rank = input_tensor_shape.dims();
const absl::InlinedVector<int64_t, 4> input_shape =
input_tensor_shape.dim_sizes();
const std::vector<xla::DExpr> input_exprs =
input_tensor_shape.get_filled_expressions();
const int block_rank = block_shape.size();

OP_REQUIRES(
Expand Down Expand Up @@ -76,11 +78,21 @@ void BatchToSpace(XlaOpKernelContext* ctx, const xla::XlaOp input,
") is not divisible by product of block sizes (",
block_num_elems, ")"));
std::vector<int64_t> reshaped_shape(input_rank + block_rank);
std::vector<xla::DExpr> reshaped_exprs(input_rank + block_rank);
std::copy(block_shape.begin(), block_shape.end(), reshaped_shape.begin());
std::fill(reshaped_exprs.begin(), reshaped_exprs.begin() + block_rank,
xla::DExpr::Const(0));
for (int i = 0; i < block_rank; ++i) {
reshaped_exprs[i] = xla::DExpr::Const(block_shape[i]);
}
reshaped_shape[block_rank] = batch_size / block_num_elems;
reshaped_exprs[block_rank] =
(input_exprs[0] / xla::DExpr::Const(block_num_elems)).simplify();
std::copy(input_shape.begin() + 1, input_shape.end(),
reshaped_shape.begin() + block_rank + 1);
xla::XlaOp reshaped = xla::Reshape(input, reshaped_shape);
std::copy(input_exprs.begin() + 1, input_exprs.end(),
reshaped_exprs.begin() + block_rank + 1);
xla::XlaOp reshaped = xla::Reshape(input, reshaped_shape, reshaped_exprs);

// 2. Permute dimensions of `reshaped` to produce `permuted` of shape
// [batch / prod(block_shape),
Expand Down Expand Up @@ -111,15 +123,22 @@ void BatchToSpace(XlaOpKernelContext* ctx, const xla::XlaOp input,
// ...,
// input_shape[N-1]]
std::vector<int64_t> reshaped_permuted_shape(input_rank);
std::vector<xla::DExpr> reshaped_permuted_exprs(input_rank);
reshaped_permuted_shape[0] = batch_size / block_num_elems;
reshaped_permuted_exprs[0] =
(input_exprs[0] / xla::DExpr::Const(block_num_elems)).simplify();
for (int i = 0; i < block_rank; ++i) {
reshaped_permuted_shape[1 + i] = block_shape[i] * input_shape[1 + i];
reshaped_permuted_exprs[1 + i] =
(xla::DExpr::Const(block_shape[i]) * input_exprs[1 + i]).simplify();
}
std::copy(remainder_shape.begin(), remainder_shape.end(),
reshaped_permuted_shape.begin() + 1 + block_rank);
std::copy(input_exprs.begin() + 1 + block_rank, input_exprs.end(),
reshaped_permuted_exprs.begin() + 1 + block_rank);

xla::XlaOp reshaped_permuted =
xla::Reshape(permuted, reshaped_permuted_shape);
xla::Reshape(permuted, reshaped_permuted_shape, reshaped_permuted_exprs);

// 4. Crop the start and end of dimensions `[1, ..., M]` of
// `reshaped_permuted` according to `crops` to produce the output of shape:
Expand All @@ -133,21 +152,26 @@ void BatchToSpace(XlaOpKernelContext* ctx, const xla::XlaOp input,
std::vector<int64_t> start_indices(input_rank, 0);
std::vector<int64_t> end_indices = reshaped_permuted_shape;
std::vector<int64_t> strides(input_rank, 1);
std::vector<xla::DExpr> start_exprs(input_rank, xla::DExpr::Const(0));
std::vector<xla::DExpr> end_exprs(reshaped_permuted_exprs.begin(),
reshaped_permuted_exprs.end());
for (int i = 0; i < block_rank; ++i) {
int64_t crop_start = crops.Get<int64_t>({i, 0});
int64_t crop_end = crops.Get<int64_t>({i, 1});
OP_REQUIRES(ctx, crop_start >= 0 && crop_end >= 0,
errors::InvalidArgument("Crops must be non-negative"));
start_indices[1 + i] = crop_start;
end_indices[1 + i] -= crop_end;
start_exprs[1 + i] = xla::DExpr::Const(crop_start);
end_exprs[1 + i] = (reshaped_permuted_exprs[1 + i] - crop_end).simplify();
OP_REQUIRES(
ctx, start_indices[1 + i] <= end_indices[1 + i],
errors::InvalidArgument(
"Cropped size must be non-negative: start: ", crop_start,
" end: ", crop_end, " size ", reshaped_permuted_shape[1 + i]));
}
xla::XlaOp output =
xla::Slice(reshaped_permuted, start_indices, end_indices, strides);
xla::XlaOp output = xla::Slice(reshaped_permuted, start_indices, end_indices,
start_exprs, end_exprs, strides);
ctx->SetOutput(0, output);
}

Expand Down
13 changes: 9 additions & 4 deletions tensorflow/compiler/tf2xla/kernels/bincount_op.cc
Original file line number Diff line number Diff line change
Expand Up @@ -110,11 +110,15 @@ class DenseBincountOp : public XlaOpKernel {
scatter_dnums.add_scatter_dims_to_operand_dims(0);

if (rank == 2) {
output_shape = xla::ShapeUtil::MakeShape(dtype, {size, output_size});
output_shape = xla::ShapeUtil::MakeShape(
dtype, {size, output_size},
std::vector<xla::DExpr>{input_shape.expressions(0),
xla::DExpr::Const(output_size)});
scatter_dnums.add_inserted_window_dims(1);
scatter_dnums.add_scatter_dims_to_operand_dims(1);
auto i_shape =
xla::ShapeUtil::MakeShape(input_xla_type, {input_shape.dimensions()});
auto i_shape = xla::ShapeUtil::MakeShape(input_xla_type,
input_shape.dimensions(),
input_shape.expressions());
auto i = xla::Iota(ctx->builder(), i_shape, 0);
xla::DExpr flattened_expr =
input_shape.expressions(0) * input_shape.expressions(1);
Expand All @@ -131,7 +135,8 @@ class DenseBincountOp : public XlaOpKernel {
updates = xla::Broadcast(
one, {input_shape.dimensions(0) * input_shape.dimensions(1)});
output = xla::Broadcast(
zero, {output_shape.dimensions(0), output_shape.dimensions(1)});
zero, {output_shape.dimensions(0), output_shape.dimensions(1)},
{output_shape.expressions(0), output_shape.expressions(1)});
if (has_weights && !binary_output_) {
weights = xla::Reshape(
weights, {input_shape.dimensions(0) * input_shape.dimensions(1)},
Expand Down
Loading