Skip to content
Merged
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
232 changes: 185 additions & 47 deletions packages/river_hdl/lib/src/core/iterative_fp_arith.dart
Original file line number Diff line number Diff line change
Expand Up @@ -36,18 +36,14 @@ import 'package:rohd/rohd.dart';
/// fnmsub (-(a*b)) + c [selFma, selNegA]
/// fnmadd (-(a*b)) + (-c) [selFma, selNegA, selNegB]
/// fcvt a, re-rounded at the other precision [selCvt]
/// A fused multiply-add runs TWO passes through the unit: pass 0 multiplies
/// a by b, pass 1 adds c to that product. Each pass rounds, so the answer is
/// the same one the golden model gives, which computes the product and the
/// sum as two separate roundings.
/// A fused multiply-add retains the complete product through alignment and
/// addition, including cancellation, and rounds only the final sum.
///
/// PRECISION. A unit wider than binary32 also serves the binary32 family, and
/// it does so without a converter on either side: [selSingle] tells it to read
/// the operand fields at binary32 and to round the answer back to binary32.
/// The widening is exact, and the answer is rounded ONCE, so it is the
/// correctly rounded binary32 answer. That matches what the golden model gives,
/// because binary64 keeps 53 bits where the bound for innocuous double
/// rounding into binary32 is 2*24 + 2 = 50.
/// Widening is exact; guard and sticky information is retained until the
/// answer is rounded once at the destination precision.
///
/// [selCvt] flips the destination precision against the source, which is what
/// fcvt.s.d and fcvt.d.s ask for. Such a convert runs as an add of the operand
Expand Down Expand Up @@ -110,7 +106,8 @@ class IterativeFpArith extends Module {
// The working exponent is signed and reaches about -(2*bias + 2*m).
final expW = e + 3;
final bias = (1 << (e - 1)) - 1;
final cntW = (ww + 2).bitLength;
final fw = 2 * n + 4; // full product, carry and guard/round/sticky
final cntW = (fw + 2).bitLength;

// The narrow format. Every internal value stays in the WIDE exponent
// domain, so a narrow operand adds this offset on the way in and a narrow
Expand Down Expand Up @@ -154,7 +151,12 @@ class IterativeFpArith extends Module {
sNorm = 7,
sDenorm = 8,
sRound = 9,
sDone = 10;
sDone = 10,
sFmaNormC = 11,
sFmaSetup = 12,
sFmaAlign = 13,
sFmaAdd = 14,
sFmaNorm = 15;
Logic st(int v) => Const(v, width: 4);

const opAdd = 0, opMul = 1, opDiv = 2;
Expand All @@ -175,7 +177,12 @@ class IterativeFpArith extends Module {
final sgn = Logic(name: 'sgn');
final effSub = Logic(name: 'effSub');
final curOp = Logic(name: 'curOp', width: 2);
final phase = Logic(name: 'phase');
final fmaProduct = Logic(name: 'fmaProduct', width: 2 * n);
final fmaExponent = Logic(name: 'fmaExponent', width: expW);
final fmaC = Logic(name: 'fmaC', width: n);
final fmaCExponent = Logic(name: 'fmaCExponent', width: expW);
final fmaLarge = Logic(name: 'fmaLarge', width: fw);
final fmaWork = Logic(name: 'fmaWork', width: fw);
// The answer takes the narrow format. Latched, because the pack stage
// reads it many cycles after the operands went in.
final dstNarrow = Logic(name: 'dstNarrow');
Expand All @@ -185,30 +192,22 @@ class IterativeFpArith extends Module {

// ---------------------------------------------------------------- decode

// Which operands this pass takes. A fused multiply-add gives pass 0 the
// raw a and b, and pass 1 the product it just made and c. The sign flips
// belong to the ADD pass, so pass 0 never applies them.
final ph0 = selFma & ~phase;
final ph1 = selFma & phase;
final negAeff = ~ph0 & ~selDiv & selNegA;
final negBeff = ~ph0 & ~selDiv & selNegB;
final rawA = mux(ph1, resReg, opA);
final rawB = mux(ph1, opC, opB);
// FMA sign controls apply to the product and addend, not the multiply's
// input operands. Keep the product unrounded until the three-operand sum.
final negAeff = ~selFma & ~selDiv & selNegA;
final negBeff = ~selFma & ~selDiv & selNegB;
final rawA = opA;
final rawB = opB;

final cvt = dual ? selCvt : Const(0);
// The product a fused multiply-add makes on pass 0 is a WIDE value, so
// pass 1 reads its left operand at the wide format whatever the operation
// precision is, and pass 0 writes a wide product for it to read.
final narrowA = dual ? (selSingle & ~ph1) : Const(0);
final narrowA = dual ? selSingle : Const(0);
final narrowB = dual ? selSingle : Const(0);
final narrowOut = dual
? (~ph0 & mux(cvt, ~selSingle, selSingle))
: Const(0);
final narrowOut = dual ? mux(cvt, ~selSingle, selSingle) : Const(0);

final opNew = mux(
selDiv,
oc(opDiv),
mux(selMul | ph0, oc(opMul), oc(opAdd)),
mux(selMul | selFma, oc(opMul), oc(opAdd)),
);
final newIsAdd = opNew.eq(oc(opAdd));
final newIsMul = opNew.eq(oc(opMul));
Expand Down Expand Up @@ -268,6 +267,7 @@ class IterativeFpArith extends Module {

final ua = unpack(rawA, narrowA, negAeff);
final ub = unpack(rawB, narrowB, negBeff);
final uc = unpack(opC, narrowB, selNegB);

// A convert reads one operand only. It runs as an add of that operand and
// a zero, so the right side is silenced here: same exponent, so nothing
Expand All @@ -283,8 +283,8 @@ class IterativeFpArith extends Module {

// An add wants the larger magnitude first, so the alignment only ever
// shifts the right operand and the subtract never borrows. The compare is
// on the unpacked exponent and significand, because a fused multiply-add
// holds its two operands at DIFFERENT formats on the second pass.
// on the unpacked exponent and significand, including the conversion's
// synthetic zero addend.
final aBigger = ua.exp.gt(bExp) | (ua.exp.eq(bExp) & ua.sig.gte(bSig));
final swapAB = newIsAdd & ~aBigger;

Expand All @@ -305,8 +305,7 @@ class IterativeFpArith extends Module {
// family costs one narrow mux and no wide one. An add with ONE zero
// operand is not special: it runs the ordinary path with a zero addend,
// which re-rounds the other operand at the destination format. That is
// what a convert needs and what a fused multiply-add needs when its third
// operand is zero and the product still has to round.
// what a convert needs. FMA handles its three-operand specials separately.
final xorSign = signA ^ signB;
final anyNaN = aNaN | bNaNs;
final mulBad = (aInf & bZero) | (bInfs & aZero);
Expand Down Expand Up @@ -352,6 +351,42 @@ class IterativeFpArith extends Module {
? mux(narrowOut, constNarrow(), constWide())
: constWide();

// Three-operand FMA special cases must be resolved before rounding. A
// finite product that exceeds the output range is not an infinity here.
final fmaProductSign = ua.sign ^ ub.sign ^ selNegA;
final fmaProductInf = (ua.isInf | ub.isInf) & ~anyNaN & ~mulBad;
final fmaInvalidSum = fmaProductInf & uc.isInf & (fmaProductSign ^ uc.sign);
final fmaNaN = anyNaN | uc.isNaN | mulBad | fmaInvalidSum;
final fmaInf = ~fmaNaN & (fmaProductInf | uc.isInf);
final fmaZero = (ua.isZero | ub.isZero) & uc.isZero;
final fmaSign =
~fmaNaN &
mux(
fmaProductInf,
fmaProductSign,
mux(uc.isInf, uc.sign, fmaProductSign & uc.sign),
);
Logic fmaConstant(bool narrow) {
final eb = narrow ? eS : e;
final mb = narrow ? mS : m;
return [
if (narrow && w > wS) Const(0, width: w - wS),
fmaSign,
mux(
fmaNaN | fmaInf,
Const((1 << eb) - 1, width: eb),
Const(0, width: eb),
),
fmaNaN,
Const(0, width: mb - 1),
].swizzle();
}

final fmaSpecialValue = dual
? mux(narrowOut, fmaConstant(true), fmaConstant(false))
: fmaConstant(false);
final initialSpecial = mux(selFma, fmaNaN | fmaInf | fmaZero, specHit);

// ------------------------------------------------------------- iteration

// Sticky-preserving right shift: bit 0 keeps the OR of everything that has
Expand Down Expand Up @@ -415,6 +450,34 @@ class IterativeFpArith extends Module {

final loadIsMul = curOp.eq(oc(opMul));

final fmaPWide = [Const(0), fmaProduct, Const(0, width: 3)].swizzle();
final fmaCWide = [Const(0), fmaC, Const(0, width: n + 3)].swizzle();
final productExpGreater = mux(
fmaExponent[expW - 1] ^ fmaCExponent[expW - 1],
~fmaExponent[expW - 1],
fmaExponent.gt(fmaCExponent),
);
final productBigger =
~fmaC.or() |
(fmaProduct.or() &
(productExpGreater |
(fmaExponent.eq(fmaCExponent) & fmaPWide.gte(fmaCWide))));
final fmaBigExponent = mux(productBigger, fmaExponent, fmaCExponent);
final fmaSmallExponent = mux(productBigger, fmaCExponent, fmaExponent);
final fmaSmall = mux(productBigger, fmaCWide, fmaPWide);
final fmaDistance = mux(
fmaSmall.or(),
fmaBigExponent - fmaSmallExponent,
Const(0, width: expW),
);
final fmaCap = mux(
fmaDistance.gt(Const(fw, width: expW)),
Const(fw, width: expW),
fmaDistance,
);
Logic fmaShift(Logic v) =>
[Const(0), v.slice(fw - 1, 2), v[1] | v[0]].swizzle();

// Normalise. A carry out of the significand shifts right once; a cancelling
// subtract shifts left until the top bit appears.
final normCarry = rw[ww - 1];
Expand Down Expand Up @@ -523,7 +586,12 @@ class IterativeFpArith extends Module {
sgn < 0,
effSub < 0,
curOp < 0,
phase < 0,
fmaProduct < 0,
fmaExponent < 0,
fmaC < 0,
fmaCExponent < 0,
fmaLarge < 0,
fmaWork < 0,
dstNarrow < 0,
special < 0,
specialHold < 0,
Expand All @@ -543,9 +611,15 @@ class IterativeFpArith extends Module {
expA < expAload,
expB < expBload,
dstNarrow < narrowOut,
special < specHit,
specialHold < specVal,
If(specHit, then: [state < sRound], orElse: [state < sNormA]),
fmaC < uc.sig,
fmaCExponent < uc.exp,
special < initialSpecial,
specialHold < mux(selFma, fmaSpecialValue, specVal),
If(
initialSpecial,
then: [state < sRound],
orElse: [state < sNormA],
),
],
),
]),
Expand Down Expand Up @@ -610,9 +684,79 @@ class IterativeFpArith extends Module {
]),
CaseItem(st(sAddSub), [rw < addRes, state < sNorm]),
CaseItem(st(sLoad), [
rw < mux(loadIsMul, mulR, divR),
expR < mux(loadIsMul, mulExp, divExp),
state < sNorm,
If(
selFma,
then: [
fmaProduct <
mux(
prodHi,
prod,
[prod.slice(2 * n - 2, 0), Const(0)].swizzle(),
),
fmaExponent < mulExp,
state < sFmaNormC,
],
orElse: [
rw < mux(loadIsMul, mulR, divR),
expR < mux(loadIsMul, mulExp, divExp),
state < sNorm,
],
),
]),
CaseItem(st(sFmaNormC), [
If(
fmaC[n - 1] | ~fmaC.or(),
then: [state < sFmaSetup],
orElse: [
fmaC < [fmaC.slice(n - 2, 0), Const(0)].swizzle(),
fmaCExponent < fmaCExponent - 1,
],
),
]),
CaseItem(st(sFmaSetup), [
fmaLarge < mux(productBigger, fmaPWide, fmaCWide),
fmaWork < fmaSmall,
expR < fmaBigExponent,
sgn < mux(productBigger, fmaProductSign, uc.sign),
effSub < (fmaProductSign ^ uc.sign),
cnt < fmaCap.getRange(0, cntW),
state < sFmaAlign,
]),
CaseItem(st(sFmaAlign), [
If(
cnt.eq(0),
then: [state < sFmaAdd],
orElse: [fmaWork < fmaShift(fmaWork), cnt < cnt - 1],
),
]),
CaseItem(st(sFmaAdd), [
fmaWork < mux(effSub, fmaLarge - fmaWork, fmaLarge + fmaWork),
state < sFmaNorm,
]),
CaseItem(st(sFmaNorm), [
If(
fmaWork[fw - 1],
then: [fmaWork < fmaShift(fmaWork), expR < expR + 1],
orElse: [
If(
fmaWork[fw - 2] | ~fmaWork.or(),
then: [
// Jam only AFTER cancellation: no product bit may be lost
// before its position in the final significand is known.
rw <
[
fmaWork.slice(fw - 1, n + 1),
fmaWork.slice(n, 0).or(),
].swizzle(),
state < sNorm,
],
orElse: [
fmaWork < [fmaWork.slice(fw - 2, 0), Const(0)].swizzle(),
expR < expR - 1,
],
),
],
),
]),
CaseItem(st(sNorm), [
If(
Expand Down Expand Up @@ -660,16 +804,10 @@ class IterativeFpArith extends Module {
]),
CaseItem(st(sRound), [
resReg < mux(special, specialHold, packed),
// A fused multiply-add takes a second pass: the product goes
// back in as the left operand and c comes in as the right one.
If(
selFma & ~phase,
then: [phase < 1, state < sIdle],
orElse: [state < sDone],
),
state < sDone,
]),
CaseItem(st(sDone), [
If(~start, then: [state < sIdle, phase < 0]),
If(~start, then: [state < sIdle]),
]),
]),
],
Expand Down
Loading
Loading