CIRCT 24.0.0git
Loading...
Searching...
No Matches
DatapathToComb.cpp
Go to the documentation of this file.
1//===----------------------------------------------------------------------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8
14#include "mlir/Analysis/TopologicalSortUtils.h"
15#include "mlir/Dialect/Func/IR/FuncOps.h"
16#include "mlir/IR/PatternMatch.h"
17#include "mlir/Pass/Pass.h"
18#include "mlir/Transforms/GreedyPatternRewriteDriver.h"
19#include "llvm/Support/Debug.h"
20#include "llvm/Support/KnownBits.h"
21#include <algorithm>
22
23#define DEBUG_TYPE "datapath-to-comb"
24
25namespace circt {
26#define GEN_PASS_DEF_CONVERTDATAPATHTOCOMB
27#include "circt/Conversion/Passes.h.inc"
28} // namespace circt
29
30using namespace circt;
31using namespace datapath;
32
33// A wrapper for comb::extractBits that returns a SmallVector<Value>.
34static SmallVector<Value> extractBits(OpBuilder &builder, Value val) {
35 SmallVector<Value> bits;
36 comb::extractBits(builder, val, bits);
37 return bits;
38}
39
40// Check whether a value is zero-extended or sign-extended - and return the
41// unextended base value and whether it was sign-extended.
42static std::pair<bool, Value> getBaseOfExt(PatternRewriter &rewriter,
43 Location loc, Value val) {
44
45 Value replBits;
46 // Check for zext
47 if (matchPattern(val, comb::m_ZextBy(mlir::matchers::m_Any(&replBits)))) {
48 auto baseWidth = val.getType().getIntOrFloatBitWidth() -
49 replBits.getType().getIntOrFloatBitWidth();
50 auto valBase =
51 rewriter.createOrFold<comb::ExtractOp>(loc, val, 0, baseWidth);
52 return {false, valBase};
53 }
54
55 // Check for sext of the value
56 if (matchPattern(val, comb::m_SextBy(mlir::matchers::m_Any(&replBits)))) {
57 auto baseWidth = val.getType().getIntOrFloatBitWidth() -
58 replBits.getType().getIntOrFloatBitWidth();
59 auto valBase =
60 rewriter.createOrFold<comb::ExtractOp>(loc, val, 0, baseWidth);
61 return {true, valBase};
62 }
63
64 // Not extended, return original value
65 return {false, val};
66}
67
68//===----------------------------------------------------------------------===//
69// Conversion patterns
70//===----------------------------------------------------------------------===//
71
72namespace {
73// Replace compressor by an adder of the inputs and zero for the other results:
74// compress(a,b,c,d) -> {a+b+c+d, 0}
75// Facilitates use of downstream compression algorithms e.g. Yosys
76struct DatapathCompressOpAddConversion : mlir::OpRewritePattern<CompressOp> {
78 LogicalResult
79 matchAndRewrite(CompressOp op,
80 mlir::PatternRewriter &rewriter) const override {
81 Location loc = op.getLoc();
82 auto inputs = op.getOperands();
83 unsigned width = inputs[0].getType().getIntOrFloatBitWidth();
84 // Sum all the inputs - set that to result value 0
85 auto addOp = comb::AddOp::create(rewriter, loc, inputs, true);
86 // Replace remaining results with zeros
87 auto zeroOp = hw::ConstantOp::create(rewriter, loc, APInt(width, 0));
88 SmallVector<Value> results(op.getNumResults() - 1, zeroOp);
89 results.push_back(addOp);
90 rewriter.replaceOp(op, results);
91 return success();
92 }
93};
94
95// Replace compressor by a wallace tree of full-adders
96struct DatapathCompressOpConversion : mlir::OpRewritePattern<CompressOp> {
97 DatapathCompressOpConversion(MLIRContext *context,
99 : mlir::OpRewritePattern<CompressOp>(context), analysis(analysis) {}
100
101 LogicalResult
102 matchAndRewrite(CompressOp op,
103 mlir::PatternRewriter &rewriter) const override {
104 Location loc = op.getLoc();
105 auto inputs = op.getOperands();
106
107 SmallVector<SmallVector<Value>> addends;
108 for (auto input : inputs) {
109 addends.push_back(
110 extractBits(rewriter, input)); // Extract bits from each input
111 }
112
113 // Compressor tree reduction
114 auto width = inputs[0].getType().getIntOrFloatBitWidth();
115 auto targetAddends = op.getNumResults();
116 datapath::CompressorTree comp(width, addends, loc, rewriter);
117
118 if (analysis) {
119 // Update delay information with arrival times
120 if (failed(comp.withInputDelays(
121 [&](Value v) { return analysis->getMaxDelay(v, 0); })))
122 return failure();
123 }
124
125 rewriter.replaceOp(op, comp.compressToHeight(rewriter, targetAddends));
126 return success();
127 }
128
129private:
130 synth::IncrementalLongestPathAnalysis *analysis = nullptr;
131};
132
133struct DatapathPartialProductOpConversion : OpRewritePattern<PartialProductOp> {
134 using OpRewritePattern<PartialProductOp>::OpRewritePattern;
135
136 DatapathPartialProductOpConversion(MLIRContext *context, bool forceBooth)
137 : OpRewritePattern<PartialProductOp>(context), forceBooth(forceBooth){};
138
139 const bool forceBooth;
140
141 LogicalResult matchAndRewrite(PartialProductOp op,
142 PatternRewriter &rewriter) const override {
143
144 Value a = op.getLhs();
145 Value b = op.getRhs();
146 unsigned width = a.getType().getIntOrFloatBitWidth();
147
148 // Skip a zero width value.
149 if (width == 0) {
150 rewriter.replaceOpWithNewOp<hw::ConstantOp>(op, op.getType(0), 0);
151 return success();
152 }
153
154 // Square partial product array can be reduced to upper triangular array.
155 // For example: AND array for a 4-bit squarer:
156 // 0 0 0 a0a3 a0a2 a0a1 a0a0
157 // 0 0 a1a3 a1a2 a1a1 a1a0 0
158 // 0 a2a3 a2a2 a2a1 a2a0 0 0
159 // a3a3 a3a2 a3a1 a3a0 0 0 0
160 //
161 // Can be reduced to:
162 // 0 0 a0a3 a0a2 a0a1 0 a0
163 // 0 a1a3 a1a2 0 a1 0 0
164 // a2a3 0 a2 0 0 0 0
165 // a3 0 0 0 0 0 0
166 if (a == b)
167 return lowerSqrAndArray(rewriter, a, op, width);
168
169 // Use result rows as a heuristic to guide partial product
170 // implementation
171 if (comb::shouldUseBoothEncoding(a, b) || forceBooth)
172 return lowerBoothArray(rewriter, a, b, op, width);
173 else
174 return lowerAndArray(rewriter, a, b, op, width);
175 }
176
177private:
178 static LogicalResult lowerAndArray(PatternRewriter &rewriter, Value a,
179 Value b, PartialProductOp op,
180 unsigned width) {
181
182 Location loc = op.getLoc();
183 // Keep a as a bitvector - multiply by each digit of b
184 SmallVector<Value> bBits = extractBits(rewriter, b);
185 auto knownBitsB = comb::computeKnownBits(b);
186
187 auto rowWidth = width;
188 auto knownBitsA = comb::computeKnownBits(a);
189 if (!knownBitsA.Zero.isZero()) {
190 if (knownBitsA.Zero.countLeadingOnes() > 1) {
191 rowWidth -= knownBitsA.Zero.countLeadingOnes();
192 a = rewriter.createOrFold<comb::ExtractOp>(loc, a, 0, rowWidth);
193 }
194 }
195
196 SmallVector<Value> partialProducts;
197 partialProducts.reserve(width);
198 // AND Array Construction:
199 // partialProducts[i] = ({b[i],..., b[i]} & a) << i
200 assert(op.getNumResults() <= width &&
201 "Cannot return more results than the operator width");
202
203 for (unsigned i = 0; i < op.getNumResults(); ++i) {
204 // Constuct partial product row for bit i of b.
205 Value ppRow;
206
207 // Skip generation of zero rows.
208 if (knownBitsB.Zero[i]) {
209 partialProducts.push_back(
210 hw::ConstantOp::create(rewriter, loc, APInt(width, 0)));
211 continue;
212 }
213
214 // If the bit is known to be one, just use `a` as the partial product row.
215 if (knownBitsB.One[i]) {
216 ppRow = a;
217 } else {
218 auto repl =
219 rewriter.createOrFold<comb::ReplicateOp>(loc, bBits[i], rowWidth);
220 ppRow = rewriter.createOrFold<comb::AndOp>(loc, repl, a);
221 }
222 if (rowWidth < width) {
223 auto padding = width - rowWidth;
224 auto zeroPad = hw::ConstantOp::create(rewriter, loc, APInt(padding, 0));
225 ppRow = rewriter.createOrFold<comb::ConcatOp>(
226 loc, ValueRange{zeroPad, ppRow}); // Pad to full width
227 }
228
229 if (i == 0) {
230 partialProducts.push_back(ppRow);
231 continue;
232 }
233 auto shiftBy = hw::ConstantOp::create(rewriter, loc, APInt(i, 0));
234 auto ppAlign =
235 comb::ConcatOp::create(rewriter, loc, ValueRange{ppRow, shiftBy});
236 auto ppAlignTrunc = rewriter.createOrFold<comb::ExtractOp>(
237 loc, ppAlign, 0, width); // Truncate to width+i bits
238 partialProducts.push_back(ppAlignTrunc);
239 }
240
241 rewriter.replaceOp(op, partialProducts);
242 return success();
243 }
244
245 static LogicalResult lowerSqrAndArray(PatternRewriter &rewriter, Value a,
246 PartialProductOp op, unsigned width) {
247
248 Location loc = op.getLoc();
249 SmallVector<Value> aBits = extractBits(rewriter, a);
250
251 SmallVector<Value> partialProducts;
252 partialProducts.reserve(width);
253 // AND Array Construction - reducing to upper triangle:
254 // partialProducts[i] = ({a[i],..., a[i]} & a) << i
255 // optimised to: {a[i] & a[n-1], ..., a[i] & a[i+1], 0, a[i], 0, ..., 0}
256 assert(op.getNumResults() <= width &&
257 "Cannot return more results than the operator width");
258 auto zeroFalse = hw::ConstantOp::create(rewriter, loc, APInt(1, 0));
259 for (unsigned i = 0; i < op.getNumResults(); ++i) {
260 SmallVector<Value> row;
261 row.reserve(width);
262
263 if (2 * i >= width) {
264 // Pad the remaining rows with zeros
265 auto zeroWidth = hw::ConstantOp::create(rewriter, loc, APInt(width, 0));
266 partialProducts.push_back(zeroWidth);
267 continue;
268 }
269
270 if (i > 0) {
271 auto shiftBy = hw::ConstantOp::create(rewriter, loc, APInt(2 * i, 0));
272 row.push_back(shiftBy);
273 }
274 row.push_back(aBits[i]);
275
276 // Track width of constructed row
277 unsigned rowWidth = 2 * i + 1;
278 if (rowWidth < width) {
279 row.push_back(zeroFalse);
280 ++rowWidth;
281 }
282
283 for (unsigned j = i + 1; j < width; ++j) {
284 // Stop when we reach the required width
285 if (rowWidth == width)
286 break;
287
288 // Otherwise pad with zeros or partial product bits
289 ++rowWidth;
290 // Number of results indicates number of non-zero bits in input
291 if (j >= op.getNumResults()) {
292 row.push_back(zeroFalse);
293 continue;
294 }
295
296 auto ppBit =
297 rewriter.createOrFold<comb::AndOp>(loc, aBits[i], aBits[j]);
298 row.push_back(ppBit);
299 }
300 std::reverse(row.begin(), row.end());
301 auto ppRow = comb::ConcatOp::create(rewriter, loc, row);
302 partialProducts.push_back(ppRow);
303 }
304
305 rewriter.replaceOp(op, partialProducts);
306 return success();
307 }
308
309 static LogicalResult lowerBoothArray(PatternRewriter &rewriter, Value a,
310 Value b, PartialProductOp op,
311 unsigned width) {
312 // TODO: sort a and b based on non-zero bits to encode the smaller input
313 Location loc = op.getLoc();
314 auto zeroFalse = hw::ConstantOp::create(rewriter, loc, APInt(1, 0));
315
316 auto [aSigned, aBase] = getBaseOfExt(rewriter, loc, op.getLhs());
317 auto [bSigned, bBase] = getBaseOfExt(rewriter, loc, op.getRhs());
318
319 auto aBaseWidth = aBase.getType().getIntOrFloatBitWidth();
320 auto bBaseWidth = bBase.getType().getIntOrFloatBitWidth();
321
322 // Detect leading zeros in multiplicand due to zero-extension
323 // and truncate to reduce partial product bits {'0, a} * {'0, b}
324 auto rowWidth = width;
325 if (aBaseWidth < width) {
326 // Retain one leading zero/sign-bit to represent 2*a
327 rowWidth = aBaseWidth + 1;
328 a = rewriter.createOrFold<comb::ExtractOp>(loc, a, 0, rowWidth);
329 }
330
331 // Booth encoding will select each row from {-2a, -1a, 0, 1a, 2a}
332 Value twoAPre =
333 rewriter.createOrFold<comb::ConcatOp>(loc, ValueRange{a, zeroFalse});
334 Value twoA = rewriter.createOrFold<comb::ExtractOp>(
335 loc, twoAPre, 0, rowWidth); // Truncate to width bits
336
337 // Encode based on the bits of b
338
339 SmallVector<Value> bBits = extractBits(rewriter, b);
340 // Pad with two zeros - for case where there's no extensions
341 bBits.append(2, zeroFalse);
342
343 // Retain two leading zeros as when b has an even number of bits we just
344 // need to retain two leading zeros
345 if (!bSigned)
346 bBits.resize(bBaseWidth + 2);
347
348 // If b is signed, we need to sign-extend with a single sign-bit
349 if (bSigned)
350 bBits.resize(bBaseWidth + 1);
351
352 SmallVector<Value> partialProducts;
353 partialProducts.reserve(op.getNumResults());
354
355 // Booth encoding halves array height by grouping three bits at a time:
356 // partialProducts[i] = a * (-2*b[2*i+1] + b[2*i] + b[2*i-1]) << 2*i
357 // encNeg \approx (-2*b[2*i+1] + b[2*i] + b[2*i-1]) <= 0
358 // encOne = (-2*b[2*i+1] + b[2*i] + b[2*i-1]) == +/- 1
359 // encTwo = (-2*b[2*i+1] + b[2*i] + b[2*i-1]) == +/- 2
360 SmallVector<Value> encNegs;
361 Value encNegPrev;
362
363 // For even width - additional row contains the final sign correction
364 for (unsigned i = 0; i + 1 < bBits.size(); i += 2) {
365 // Get Booth bits: b[i+1], b[i], b[i-1] (b[-1] = 0)
366 Value bim1 = (i == 0) ? zeroFalse : bBits[i - 1];
367 Value bi = bBits[i];
368 Value bip1 = bBits[i + 1];
369
370 // Is the encoding zero or negative (an approximation)
371 Value encNeg = bip1;
372 encNegs.push_back(encNeg); // Store for sign-extension optimisation
373 // Is the encoding one = b[i] xor b[i-1]
374 Value encOne = rewriter.createOrFold<comb::XorOp>(loc, bi, bim1, true);
375 // Is the encoding two = (bip1 & ~bi & ~bim1) | (~bip1 & bi & bim1)
376 Value constOne = hw::ConstantOp::create(rewriter, loc, APInt(1, 1));
377 Value biInv = rewriter.createOrFold<comb::XorOp>(loc, bi, constOne, true);
378 Value bip1Inv =
379 rewriter.createOrFold<comb::XorOp>(loc, bip1, constOne, true);
380 Value bim1Inv =
381 rewriter.createOrFold<comb::XorOp>(loc, bim1, constOne, true);
382
383 Value andLeft = rewriter.createOrFold<comb::AndOp>(
384 loc, ValueRange{bip1Inv, bi, bim1}, true);
385 Value andRight = rewriter.createOrFold<comb::AndOp>(
386 loc, ValueRange{bip1, biInv, bim1Inv}, true);
387 Value encTwo =
388 rewriter.createOrFold<comb::OrOp>(loc, andLeft, andRight, true);
389
390 Value encNegRepl =
391 rewriter.createOrFold<comb::ReplicateOp>(loc, encNeg, rowWidth);
392 Value encOneRepl =
393 rewriter.createOrFold<comb::ReplicateOp>(loc, encOne, rowWidth);
394 Value encTwoRepl =
395 rewriter.createOrFold<comb::ReplicateOp>(loc, encTwo, rowWidth);
396
397 // Select between 2*a or 1*a or 0*a
398 Value selTwoA = rewriter.createOrFold<comb::AndOp>(loc, encTwoRepl, twoA);
399 Value selOneA = rewriter.createOrFold<comb::AndOp>(loc, encOneRepl, a);
400 Value magA =
401 rewriter.createOrFold<comb::OrOp>(loc, selTwoA, selOneA, true);
402
403 // Conditionally invert the row
404 Value ppRow =
405 rewriter.createOrFold<comb::XorOp>(loc, magA, encNegRepl, true);
406
407 // No sign-correction in the first row
408 if (i == 0) {
409 partialProducts.push_back(ppRow);
410 encNegPrev = encNeg;
411 continue;
412 }
413
414 if (i == 2) {
415 Value withSignCorrection = rewriter.createOrFold<comb::ConcatOp>(
416 loc, ValueRange{ppRow, zeroFalse, encNegPrev});
417 partialProducts.push_back(withSignCorrection);
418 encNegPrev = encNeg;
419 continue;
420 }
421
422 // Insert a sign-correction from the previous row
423 // {ppRow, 0, encNegPrev} << (i-2)
424 Value shiftBy = hw::ConstantOp::create(rewriter, loc, APInt(i - 2, 0));
425 Value withSignCorrection = rewriter.createOrFold<comb::ConcatOp>(
426 loc, ValueRange{ppRow, zeroFalse, encNegPrev, shiftBy});
427 partialProducts.push_back(withSignCorrection);
428 encNegPrev = encNeg;
429
430 if (partialProducts.size() == op.getNumResults())
431 break;
432 }
433
434 // Add the final sign-correction row for signed multiplication
435 // Not necessary for unsigned multiplication as the final row is positive
436 if (bSigned) {
437 auto numPP = partialProducts.size();
438 Value shiftByFinal =
439 hw::ConstantOp::create(rewriter, loc, APInt((numPP - 1) * 2, 0));
440 Value finalSignCorrection = rewriter.createOrFold<comb::ConcatOp>(
441 loc, ValueRange{zeroFalse, encNegPrev, shiftByFinal});
442 partialProducts.push_back(finalSignCorrection);
443 encNegs.push_back(zeroFalse); // No sign-extension for the final row
444 }
445
446 // Sign-extension:
447 // { s1, s1, s1, s1, s1, p1}
448 // { s2, s2, s2, p2 }
449 // { s3, p3 }
450 // TODO: optimize by only replicating the sign bit once using
451 // typical sign-extension trick - can be handled by separate
452 // canonicalization patterns
453 for (unsigned i = 0; i < partialProducts.size(); ++i) {
454 auto ppRow = partialProducts[i];
455
456 auto ppWidth = ppRow.getType().getIntOrFloatBitWidth();
457 if (ppWidth < width) {
458 auto padding = width - ppWidth;
459 auto encNeg = encNegs[i];
460 if (aSigned)
461 encNeg = rewriter.createOrFold<comb::ExtractOp>(loc, ppRow,
462 ppWidth - 1, 1);
463
464 // Replicate the encNeg bit for sign-extension
465 Value encNegPad =
466 rewriter.createOrFold<comb::ReplicateOp>(loc, encNeg, padding);
467 ppRow = rewriter.createOrFold<comb::ConcatOp>(
468 loc, ValueRange{encNegPad, ppRow}); // Pad to full width
469 }
470
471 // Truncate any excess bits
472 ppWidth = ppRow.getType().getIntOrFloatBitWidth();
473 if (ppWidth > width) {
474 ppRow = rewriter.createOrFold<comb::ExtractOp>(loc, ppRow, 0, width);
475 }
476 partialProducts[i] = ppRow;
477 assert(partialProducts[i].getType().getIntOrFloatBitWidth() == width &&
478 "Expected sign-extended partial product to be full width");
479 }
480
481 // Zero-pad to match the required output width
482 auto zeroWidth = hw::ConstantOp::create(rewriter, loc, APInt(width, 0));
483 while (partialProducts.size() < op.getNumResults())
484 partialProducts.push_back(zeroWidth);
485
486 assert(partialProducts.size() == op.getNumResults() &&
487 "Expected number of booth partial products to match results");
488
489 rewriter.replaceOp(op, partialProducts);
490 return success();
491 }
492};
493
494struct DatapathPosPartialProductOpConversion
495 : OpRewritePattern<PosPartialProductOp> {
496 using OpRewritePattern<PosPartialProductOp>::OpRewritePattern;
497
498 DatapathPosPartialProductOpConversion(MLIRContext *context, bool forceBooth)
499 : OpRewritePattern<PosPartialProductOp>(context),
500 forceBooth(forceBooth){};
501
502 const bool forceBooth;
503
504 LogicalResult matchAndRewrite(PosPartialProductOp op,
505 PatternRewriter &rewriter) const override {
506
507 Value a = op.getAddend0();
508 Value b = op.getAddend1();
509 Value c = op.getMultiplicand();
510 unsigned width = a.getType().getIntOrFloatBitWidth();
511
512 // Skip a zero width value.
513 if (width == 0) {
514 rewriter.replaceOpWithNewOp<hw::ConstantOp>(op, op.getType(0), 0);
515 return success();
516 }
517
518 // A one-bit product has no spare row for the correction of a negative
519 // redundant Booth digit. The direct AND-array form is already minimal.
520 if (width == 1)
521 return lowerAndArray(rewriter, a, b, c, op, width);
522
523 // Recoding the carry-save operand directly avoids materializing a
524 // carry-propagate addition before the multiplier.
525 if (comb::shouldUseBoothEncoding(a, b) || forceBooth)
526 return lowerBoothArray(rewriter, a, b, c, op, width);
527 return lowerAndArray(rewriter, a, b, c, op, width);
528 }
529
530private:
531 static LogicalResult lowerBoothArray(PatternRewriter &rewriter, Value a,
532 Value b, Value c, PosPartialProductOp op,
533 unsigned width) {
534 Location loc = op.getLoc();
535 Value zero = hw::ConstantOp::create(rewriter, loc, APInt(1, 0));
536
537 // An implementation based on Zimmerman's 2003 paper:
538 // "Optimized Synthesis of Sum-of-Products"
539 // The idea is to modify the Booth encoding to consumer 6-bits for each row
540 // of the partial product array (3 from a and 3 from b).
541 // Change from a traditional Booth encoder is that there is a carry-chain
542 // booth[i] = -2*(a[i+1]+b[i+1]) + (a[i]+b[i]) + (a[i-1]+b[i-1])
543 // + carry[i] - 4*carry[i+1]
544 //
545 // The additional carry-bits are there to keep the booth digit in the range
546 // [-2, 2] but do not constitute a full-carry chain as c[i+1] is independent
547 // of c[i].
548
549 // Handle sign/zero-extended multiplicand c
550 auto [cSigned, cBase] = getBaseOfExt(rewriter, loc, c);
551 unsigned cBaseWidth = cBase.getType().getIntOrFloatBitWidth();
552 unsigned rowWidth = width;
553 if (cBaseWidth < width) {
554 // Retain a leading zero or sign bit so that 2*c is representable.
555 rowWidth = cBaseWidth + 1;
556 c = rewriter.createOrFold<comb::ExtractOp>(loc, c, 0, rowWidth);
557 }
558
559 Value twoCPre =
560 rewriter.createOrFold<comb::ConcatOp>(loc, ValueRange{c, zero});
561 Value twoC =
562 rewriter.createOrFold<comb::ExtractOp>(loc, twoCPre, 0, rowWidth);
563
564 // Now compute the bits of the encoding pair of values
565 auto [aSigned, aBase] = getBaseOfExt(rewriter, loc, a);
566 auto [bSigned, bBase] = getBaseOfExt(rewriter, loc, b);
567
568 auto aBaseWidth = aBase.getType().getIntOrFloatBitWidth();
569 auto bBaseWidth = bBase.getType().getIntOrFloatBitWidth();
570
571 auto encodeSigned = aSigned && bSigned;
572 auto encodeBaseWidth = std::max(aBaseWidth, bBaseWidth);
573
574 SmallVector<Value> aBits = extractBits(rewriter, a);
575 SmallVector<Value> bBits = extractBits(rewriter, b);
576 // First reduce to their base widths - clip leading zeros/sign-bits
577 aBits.resize(encodeBaseWidth);
578 bBits.resize(encodeBaseWidth);
579
580 // For unsigned pad with three leading zeros
581 if (!encodeSigned) {
582 aBits.append(3, zero);
583 bBits.append(3, zero);
584 }
585
586 // If a & b are signed, we need to sign-extend by two bits
587 if (encodeSigned) {
588 aBits.append(2, aBits.back());
589 bBits.append(2, bBits.back());
590 }
591 SmallVector<Value> partialProducts;
592 SmallVector<Value> encNegs;
593 partialProducts.reserve(op.getNumResults());
594 encNegs.reserve(aBits.size());
595 Value encNegPrev;
596 Value recoderCarry = zero;
597
598 for (unsigned i = 0; i + 1 < aBits.size(); i += 2) {
599 // Select the Booth bits (first row will have b[-1] = a[-1] = 0)
600 Value aim1 = (i == 0) ? zero : aBits[i - 1];
601 Value bim1 = (i == 0) ? zero : bBits[i - 1];
602 Value ai = aBits[i];
603 Value bi = bBits[i];
604 Value aip1 = aBits[i + 1];
605 Value bip1 = bBits[i + 1];
606
607 // The implementation is entirely based on Figure 3 of
608 // "Optimized Synthesis of Sum-of-Products"
609 // which provides little intuition behind the encoding circuit - but it is
610 // really just compact logical expressions to determine the value of
611 // -2*(a[i+1]+b[i+1]) + (a[i]+b[i]) + (a[i-1]+b[i-1])
612 // + carry[i] - 4*carry[i+1]
613
614 // Compute a majority function of a[i], b[i] and b[i-1] indicating
615 // a[i] + b[i] + b[i-1] >= 2
616 Value aAndB = rewriter.createOrFold<comb::AndOp>(loc, ai, bi, true);
617 Value aAndPrevB = rewriter.createOrFold<comb::AndOp>(loc, ai, bim1, true);
618 Value bAndPrevB = rewriter.createOrFold<comb::AndOp>(loc, bi, bim1, true);
619 Value majority = rewriter.createOrFold<comb::OrOp>(
620 loc, ValueRange{aAndB, aAndPrevB, bAndPrevB}, true);
621
622 Value nextXor = rewriter.createOrFold<comb::XorOp>(loc, aip1, bip1, true);
623 Value aXorB = rewriter.createOrFold<comb::XorOp>(loc, ai, bi, true);
624 Value aXorPrevA = rewriter.createOrFold<comb::XorOp>(loc, ai, aim1, true);
625 Value prevOr = rewriter.createOrFold<comb::OrOp>(loc, aim1, bim1, true);
626
627 // Second layer of logic
628 // y1 = (a[i] ^ b[i]) ^ (a[i-1] | b[i-1])
629 Value y1 = rewriter.createOrFold<comb::XorOp>(loc, aXorB, prevOr, true);
630 Value y2 = rewriter.createOrFold<comb::OrOp>(loc, aXorB, aXorPrevA, true);
631 // z1 = (a[i+1] ^ b[i+1]) ^ ((a[i] ^ b[i]) | (a[i-1] ^ b[i-1]))
632 Value z1 = rewriter.createOrFold<comb::XorOp>(loc, nextXor, y2, true);
633
634 // Encoding signals
635 // encNeg selects a negative partial product
636 Value encNeg =
637 rewriter.createOrFold<comb::XorOp>(loc, majority, nextXor, true);
638 Value invNextXor = comb::createOrFoldNot(rewriter, loc, nextXor);
639 // Compute the carry for the next row - this is to keep the digits within
640 // the range [-2,2]
641 Value recoderCarryNext =
642 rewriter.createOrFold<comb::AndOp>(loc, invNextXor, majority, true);
643
644 // encOne selects a partial product of magnitude 1
645 Value encOne = rewriter.createOrFold<comb::XorOp>(
646 loc, ValueRange{y1, recoderCarry}, true);
647 Value encOneInv = comb::createOrFoldNot(rewriter, loc, encOne);
648 // encTwo selects a partial product of magnitude 2
649 Value encTwo = rewriter.createOrFold<comb::AndOp>(
650 loc, ValueRange{z1, encOneInv}, true);
651
652 // From here this is conventional Booth encoding!
653 Value encNegRepl =
654 rewriter.createOrFold<comb::ReplicateOp>(loc, encNeg, rowWidth);
655 Value encOneRepl =
656 rewriter.createOrFold<comb::ReplicateOp>(loc, encOne, rowWidth);
657 Value encTwoRepl =
658 rewriter.createOrFold<comb::ReplicateOp>(loc, encTwo, rowWidth);
659 Value selTwoC = rewriter.createOrFold<comb::AndOp>(loc, encTwoRepl, twoC);
660 Value selOneC = rewriter.createOrFold<comb::AndOp>(loc, encOneRepl, c);
661 Value magnitude = rewriter.createOrFold<comb::OrOp>(
662 loc, ValueRange{selTwoC, selOneC}, true);
663 Value ppRow =
664 rewriter.createOrFold<comb::XorOp>(loc, magnitude, encNegRepl, true);
665
666 encNegs.push_back(encNeg);
667 if (i == 0)
668 partialProducts.push_back(ppRow);
669 else if (i == 2)
670 partialProducts.push_back(rewriter.createOrFold<comb::ConcatOp>(
671 loc, ValueRange{ppRow, zero, encNegPrev}));
672 else {
673 Value shift = hw::ConstantOp::create(rewriter, loc, APInt(i - 2, 0));
674 partialProducts.push_back(rewriter.createOrFold<comb::ConcatOp>(
675 loc, ValueRange{ppRow, zero, encNegPrev, shift}));
676 }
677 encNegPrev = encNeg;
678 recoderCarry = recoderCarryNext;
679
680 if (partialProducts.size() == op.getNumResults())
681 break;
682 }
683
684 // Add the final sign-correction row for signed multiplication
685 // Not necessary for unsigned multiplication as the final row is positive
686 // The final recoderCarry will always be zero by construction
687 if (encodeSigned) {
688 auto numPP = partialProducts.size();
689 Value shiftByFinal =
690 hw::ConstantOp::create(rewriter, loc, APInt((numPP - 1) * 2, 0));
691 Value finalSignCorrection = rewriter.createOrFold<comb::ConcatOp>(
692 loc, ValueRange{zero, encNegPrev, shiftByFinal});
693 partialProducts.push_back(finalSignCorrection);
694 encNegs.push_back(zero); // No sign-extension for the final row
695 }
696
697 // Handle sign-extesion of the partial products to full width
698 for (auto [index, ppRow] : llvm::enumerate(partialProducts)) {
699 unsigned ppWidth = ppRow.getType().getIntOrFloatBitWidth();
700 if (ppWidth < width) {
701 Value sign = encNegs[index];
702 if (cSigned)
703 sign = rewriter.createOrFold<comb::ExtractOp>(loc, ppRow, ppWidth - 1,
704 1);
705 Value padding = rewriter.createOrFold<comb::ReplicateOp>(
706 loc, sign, width - ppWidth);
707 ppRow = rewriter.createOrFold<comb::ConcatOp>(
708 loc, ValueRange{padding, ppRow});
709 }
710 if (ppRow.getType().getIntOrFloatBitWidth() > width)
711 ppRow = rewriter.createOrFold<comb::ExtractOp>(loc, ppRow, 0, width);
712 partialProducts[index] = ppRow;
713 }
714
715 Value zeroWidth = hw::ConstantOp::create(rewriter, loc, APInt(width, 0));
716 partialProducts.resize(op.getNumResults(), zeroWidth);
717 rewriter.replaceOp(op, partialProducts);
718 return success();
719 }
720
721 static LogicalResult lowerAndArray(PatternRewriter &rewriter, Value a,
722 Value b, Value c, PosPartialProductOp op,
723 unsigned width) {
724
725 Location loc = op.getLoc();
726 // Encode (a+b) by implementing a half-adder - then note the following
727 // fact carry[i] & save[i] == false
728 auto carry = rewriter.createOrFold<comb::AndOp>(loc, a, b);
729 auto save = rewriter.createOrFold<comb::XorOp>(loc, a, b);
730
731 SmallVector<Value> carryBits = extractBits(rewriter, carry);
732 SmallVector<Value> saveBits = extractBits(rewriter, save);
733
734 // Reduce c width based on leading zeros
735 auto rowWidth = width;
736 auto [cSigned, cBase] = getBaseOfExt(rewriter, loc, c);
737 auto cBaseWidth = cBase.getType().getIntOrFloatBitWidth();
738
739 if (cBaseWidth < width && !cSigned) {
740 // Retain one leading zero to represent 2*c
741 rowWidth = cBaseWidth + 1;
742 c = rewriter.createOrFold<comb::ExtractOp>(loc, c, 0, rowWidth);
743 }
744
745 // Compute 2*c for use in array construction
746 Value zeroFalse = hw::ConstantOp::create(rewriter, loc, APInt(1, 0));
747 Value twoCPre =
748 comb::ConcatOp::create(rewriter, loc, ValueRange{c, zeroFalse});
749 Value twoC = rewriter.createOrFold<comb::ExtractOp>(loc, twoCPre, 0,
750 rowWidth); // Truncate
751
752 // AND Array Construction:
753 // pp[i] = ( (carry[i] * (c<<1)) | (save[i] * c) ) << i
754 SmallVector<Value> partialProducts;
755 partialProducts.reserve(width);
756
757 assert(op.getNumResults() <= width &&
758 "Cannot return more results than the operator width");
759
760 for (unsigned i = 0; i < op.getNumResults(); ++i) {
761 auto replSave =
762 rewriter.createOrFold<comb::ReplicateOp>(loc, saveBits[i], rowWidth);
763 auto replCarry =
764 rewriter.createOrFold<comb::ReplicateOp>(loc, carryBits[i], rowWidth);
765
766 auto ppRowSave = rewriter.createOrFold<comb::AndOp>(loc, replSave, c);
767 auto ppRowCarry =
768 rewriter.createOrFold<comb::AndOp>(loc, replCarry, twoC);
769 auto ppRow =
770 rewriter.createOrFold<comb::OrOp>(loc, ppRowSave, ppRowCarry);
771 auto ppAlign = ppRow;
772 if (i > 0) {
773 auto shiftBy = hw::ConstantOp::create(rewriter, loc, APInt(i, 0));
774 ppAlign =
775 comb::ConcatOp::create(rewriter, loc, ValueRange{ppRow, shiftBy});
776 }
777
778 // May need to truncate shifted value
779 if (rowWidth + i > width) {
780 auto ppAlignTrunc =
781 rewriter.createOrFold<comb::ExtractOp>(loc, ppAlign, 0, width);
782 partialProducts.push_back(ppAlignTrunc);
783 continue;
784 }
785 // May need to zero pad to approriate width
786 if (rowWidth + i < width) {
787 auto extPPAlign = comb::createZExt(rewriter, loc, ppAlign, width);
788 partialProducts.push_back(extPPAlign);
789 continue;
790 }
791
792 partialProducts.push_back(ppAlign);
793 }
794
795 rewriter.replaceOp(op, partialProducts);
796 return success();
797 }
798};
799
800} // namespace
801
802//===----------------------------------------------------------------------===//
803// Convert Datapath to Comb pass
804//===----------------------------------------------------------------------===//
805
806namespace {
807struct ConvertDatapathToCombPass
808 : public impl::ConvertDatapathToCombBase<ConvertDatapathToCombPass> {
809 void runOnOperation() override;
810 using ConvertDatapathToCombBase<
811 ConvertDatapathToCombPass>::ConvertDatapathToCombBase;
812};
813} // namespace
814
816 Operation *op, RewritePatternSet &&patterns,
818 // TODO: Topologically sort the operations in the module to ensure that all
819 // dependencies are processed before their users.
820 mlir::GreedyRewriteConfig config;
821 // Set the listener to update timing information
822 // HACK: Setting max iterations to 2 to ensure that the patterns are
823 // one-shot, making sure target operations are datapath operations are
824 // replaced.
825 config.setMaxIterations(2).setListener(analysis).setUseTopDownTraversal(true);
826
827 // Apply the patterns greedily
828 if (failed(mlir::applyPatternsGreedily(op, std::move(patterns), config)))
829 return failure();
830
831 return success();
832}
833
834void ConvertDatapathToCombPass::runOnOperation() {
835 RewritePatternSet patterns(&getContext());
836
837 patterns.add<DatapathPartialProductOpConversion,
838 DatapathPosPartialProductOpConversion>(patterns.getContext(),
839 forceBooth);
840 synth::IncrementalLongestPathAnalysis *analysis = nullptr;
841 if (timingAware)
842 analysis = &getAnalysis<synth::IncrementalLongestPathAnalysis>();
843
844 if (lowerCompressToAdd)
845 // Lower compressors to simple add operations for downstream optimisations
846 patterns.add<DatapathCompressOpAddConversion>(patterns.getContext());
847 if (lowerCompress)
848 // Lower compressors to a complete gate-level implementation
849 patterns.add<DatapathCompressOpConversion>(patterns.getContext(), analysis);
850
852 getOperation(), std::move(patterns), analysis)))
853 return signalPassFailure();
854
855 // Verify that all Datapath operations have been successfully converted.
856 // Walk the operation and check for any remaining Datapath dialect
857 // operations.
858 auto result = getOperation()->walk([&](Operation *op) {
859 if (llvm::isa<datapath::CompressOp>(op) && !lowerCompress &&
860 !lowerCompressToAdd)
861 return WalkResult::advance();
862 if (llvm::isa_and_nonnull<datapath::DatapathDialect>(op->getDialect())) {
863 op->emitError("Datapath operation not converted: ") << *op;
864 return WalkResult::interrupt();
865 }
866 return WalkResult::advance();
867 });
868 if (result.wasInterrupted())
869 return signalPassFailure();
870}
assert(baseType &&"element must be base type")
static SmallVector< Value > extractBits(OpBuilder &builder, Value val)
static Value zeroPad(PatternRewriter &rewriter, Location loc, Value input, size_t targetWidth, size_t trailingZeros)
static std::pair< bool, Value > getBaseOfExt(PatternRewriter &rewriter, Location loc, Value val)
static SmallVector< Value > extractBits(OpBuilder &builder, Value val)
static LogicalResult applyPatternsGreedilyWithTimingInfo(Operation *op, RewritePatternSet &&patterns, synth::IncrementalLongestPathAnalysis *analysis)
static std::unique_ptr< Context > context
create(data_type, value)
Definition hw.py:433
The InstanceGraph op interface, see InstanceGraphInterface.td for more details.