Skip to content

Commit fc31231

Browse files
committed
perf: branchless square root implementation
1 parent dae52a4 commit fc31231

1 file changed

Lines changed: 39 additions & 90 deletions

File tree

src/Common.sol

Lines changed: 39 additions & 90 deletions
Original file line numberDiff line numberDiff line change
@@ -321,53 +321,32 @@ function exp2(uint256 x) pure returns (uint256 result) {
321321
/// @return result The index of the most significant bit as a uint256.
322322
/// @custom:smtchecker abstract-function-nondet
323323
function msb(uint256 x) pure returns (uint256 result) {
324-
// 2^128
325324
assembly ("memory-safe") {
326-
let factor := shl(7, gt(x, 0xFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFF))
327-
x := shr(factor, x)
328-
result := or(result, factor)
329-
}
330-
// 2^64
331-
assembly ("memory-safe") {
332-
let factor := shl(6, gt(x, 0xFFFFFFFFFFFFFFFF))
333-
x := shr(factor, x)
334-
result := or(result, factor)
335-
}
336-
// 2^32
337-
assembly ("memory-safe") {
338-
let factor := shl(5, gt(x, 0xFFFFFFFF))
339-
x := shr(factor, x)
340-
result := or(result, factor)
341-
}
342-
// 2^16
343-
assembly ("memory-safe") {
344-
let factor := shl(4, gt(x, 0xFFFF))
345-
x := shr(factor, x)
346-
result := or(result, factor)
347-
}
348-
// 2^8
349-
assembly ("memory-safe") {
350-
let factor := shl(3, gt(x, 0xFF))
351-
x := shr(factor, x)
352-
result := or(result, factor)
353-
}
354-
// 2^4
355-
assembly ("memory-safe") {
356-
let factor := shl(2, gt(x, 0xF))
357-
x := shr(factor, x)
358-
result := or(result, factor)
359-
}
360-
// 2^2
361-
assembly ("memory-safe") {
362-
let factor := shl(1, gt(x, 0x3))
363-
x := shr(factor, x)
364-
result := or(result, factor)
365-
}
366-
// 2^1
367-
// No need to shift x any more.
368-
assembly ("memory-safe") {
369-
let factor := gt(x, 0x1)
370-
result := or(result, factor)
325+
// msb7 = (x >= 2^128) ? 128 : 0
326+
let msb7 := shl(7, gt(x, 0xFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFF))
327+
x := shr(msb7, x)
328+
// msb6 = (x >= 2^64) ? 64 : 0
329+
let msb6 := shl(6, gt(x, 0xFFFFFFFFFFFFFFFF))
330+
x := shr(msb6, x)
331+
// msb5 = (x >= 2^32) ? 32 : 0
332+
let msb5 := shl(5, gt(x, 0xFFFFFFFF))
333+
x := shr(msb5, x)
334+
// msb4 = (x >= 2^16) ? 16 : 0
335+
let msb4 := shl(4, gt(x, 0xFFFF))
336+
x := shr(msb4, x)
337+
// msb3 = (x >= 2^8) ? 8 : 0
338+
let msb3 := shl(3, gt(x, 0xFF))
339+
x := shr(msb3, x)
340+
// msb2 = (x >= 2^4) ? 4 : 0
341+
let msb2 := shl(2, gt(x, 0xF))
342+
x := shr(msb2, x)
343+
// msb1 = (x >= 2^2) ? 2 : 0
344+
let msb1 := shl(1, gt(x, 0x3))
345+
x := shr(msb1, x)
346+
// msb0 = (x >= 2^1) ? 1 : 0
347+
let msb0 := gt(x, 0x1)
348+
// msb = msb7 | msb6 | msb5 | msb4 | msb3 | msb2 | msb1 | msb0
349+
result := or(or(or(or(or(or(or(msb0, msb1), msb2), msb3), msb4), msb5), msb6), msb7)
371350
}
372351
}
373352

@@ -596,10 +575,6 @@ function mulDivSigned(int256 x, int256 y, int256 denominator) pure returns (int2
596575
/// @return result The result as a uint256.
597576
/// @custom:smtchecker abstract-function-nondet
598577
function sqrt(uint256 x) pure returns (uint256 result) {
599-
if (x == 0) {
600-
return 0;
601-
}
602-
603578
// For our first guess, we calculate the biggest power of 2 which is smaller than the square root of x.
604579
//
605580
// We know that the "msb" (most significant bit) of x is a power of 2 such that we have:
@@ -623,53 +598,27 @@ function sqrt(uint256 x) pure returns (uint256 result) {
623598
// $$
624599
//
625600
// Consequently, $2^{log_2(x) /2} is a good first approximation of sqrt(x) with at least one correct bit.
626-
uint256 xAux = uint256(x);
627-
result = 1;
628-
if (xAux >= 2 ** 128) {
629-
xAux >>= 128;
630-
result <<= 64;
631-
}
632-
if (xAux >= 2 ** 64) {
633-
xAux >>= 64;
634-
result <<= 32;
635-
}
636-
if (xAux >= 2 ** 32) {
637-
xAux >>= 32;
638-
result <<= 16;
639-
}
640-
if (xAux >= 2 ** 16) {
641-
xAux >>= 16;
642-
result <<= 8;
643-
}
644-
if (xAux >= 2 ** 8) {
645-
xAux >>= 8;
646-
result <<= 4;
647-
}
648-
if (xAux >= 2 ** 4) {
649-
xAux >>= 4;
650-
result <<= 2;
651-
}
652-
if (xAux >= 2 ** 2) {
653-
result <<= 1;
601+
unchecked {
602+
// ideally, we should use arithmetic operators, but solc is not smart enough to optimize `2**(msb(x)/2)`
603+
/// forge-lint: disable-next-line(incorrect-shift)
604+
result = 1 << (msb(x) >> 1);
654605
}
655606

656607
// At this point, `result` is an estimation with at least one bit of precision. We know the true value has at
657608
// most 128 bits, since it is the square root of a uint256. Newton's method converges quadratically (precision
658609
// doubles at every iteration). We thus need at most 7 iteration to turn our partial result with one bit of
659610
// precision into the expected uint128 result.
660-
unchecked {
661-
result = (result + x / result) >> 1;
662-
result = (result + x / result) >> 1;
663-
result = (result + x / result) >> 1;
664-
result = (result + x / result) >> 1;
665-
result = (result + x / result) >> 1;
666-
result = (result + x / result) >> 1;
667-
result = (result + x / result) >> 1;
611+
assembly ("memory-safe") {
612+
// note: division by zero in EVM returns zero
613+
result := shr(1, add(result, div(x, result)))
614+
result := shr(1, add(result, div(x, result)))
615+
result := shr(1, add(result, div(x, result)))
616+
result := shr(1, add(result, div(x, result)))
617+
result := shr(1, add(result, div(x, result)))
618+
result := shr(1, add(result, div(x, result)))
619+
result := shr(1, add(result, div(x, result)))
668620

669621
// If x is not a perfect square, round the result toward zero.
670-
uint256 roundedResult = x / result;
671-
if (result >= roundedResult) {
672-
result = roundedResult;
673-
}
622+
result := sub(result, gt(result, div(x, result)))
674623
}
675624
}

0 commit comments

Comments
 (0)