Skip to content

Commit ba7259f

Browse files
authored
Merge pull request #54 from zhouguangyuan0718/codex/fp-constant-bits
ir: preserve floating-point constant bit patterns
2 parents 36244bb + cd78180 commit ba7259f

4 files changed

Lines changed: 184 additions & 0 deletions

File tree

‎IRBindings.cpp‎

Lines changed: 34 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,16 +14,50 @@
1414
#include "llvm/ADT/SmallVector.h"
1515
#include "llvm/Config/llvm-config.h"
1616
#include "llvm/IR/Attributes.h"
17+
#include "llvm/IR/Constants.h"
1718
#include "llvm/IR/DebugLoc.h"
1819
#include "llvm/IR/DebugInfoMetadata.h"
1920
#include "llvm/IR/Function.h"
2021
#include "llvm/IR/IRBuilder.h"
2122
#include "llvm/IR/Intrinsics.h"
2223
#include "llvm/IR/LLVMContext.h"
2324
#include "llvm/IR/Module.h"
25+
#include <algorithm>
2426

2527
using namespace llvm;
2628

29+
LLVMValueRef LLVMGoConstFPFromBits(LLVMTypeRef Ty, const uint64_t *Words,
30+
unsigned NumWords) {
31+
if (!Ty || !unwrap(Ty)->isFloatingPointTy())
32+
return nullptr;
33+
auto *T = unwrap(Ty);
34+
unsigned BitWidth = T->getScalarSizeInBits();
35+
if (NumWords != (BitWidth + 63) / 64 || !Words)
36+
return nullptr;
37+
if (BitWidth % 64 && (Words[NumWords - 1] >> (BitWidth % 64)))
38+
return nullptr;
39+
#if LLVM_VERSION_MAJOR >= 22
40+
return LLVMConstFPFromBits(Ty, Words);
41+
#else
42+
return wrap(ConstantFP::get(
43+
T->getContext(), APFloat(T->getFltSemantics(),
44+
APInt(BitWidth, ArrayRef<uint64_t>(Words, NumWords)))));
45+
#endif
46+
}
47+
48+
unsigned LLVMGoConstFPGetBits(LLVMValueRef Val, uint64_t *Words) {
49+
if (!Val)
50+
return 0;
51+
auto *FP = dyn_cast<ConstantFP>(unwrap(Val));
52+
if (!FP)
53+
return 0;
54+
APInt Bits = FP->getValueAPF().bitcastToAPInt();
55+
unsigned NumWords = Bits.getNumWords();
56+
if (Words)
57+
std::copy_n(Bits.getRawData(), NumWords, Words);
58+
return NumWords;
59+
}
60+
2761
LLVMAttributeRef LLVMGoCreateConstantRangeAttribute(
2862
LLVMContextRef C, unsigned KindID, unsigned NumBits,
2963
const uint64_t *LowerWords, const uint64_t *UpperWords) {

‎IRBindings.h‎

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -33,6 +33,10 @@ struct LLVMDebugLocMetadata{
3333
LLVMMetadataRef InlinedAt;
3434
};
3535

36+
LLVMValueRef LLVMGoConstFPFromBits(LLVMTypeRef Ty, const uint64_t *Words,
37+
unsigned NumWords);
38+
unsigned LLVMGoConstFPGetBits(LLVMValueRef Val, uint64_t *Words);
39+
3640
LLVMMetadataRef LLVMConstantAsMetadata(LLVMValueRef Val);
3741

3842
LLVMAttributeRef LLVMGoCreateConstantRangeAttribute(

‎float_bits_test.go‎

Lines changed: 115 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,115 @@
1+
// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
2+
// See https://llvm.org/LICENSE.txt for license information.
3+
// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
4+
5+
package llvm
6+
7+
import (
8+
"os"
9+
"path/filepath"
10+
"reflect"
11+
"testing"
12+
)
13+
14+
func TestFloatBits(t *testing.T) {
15+
// Parse independently specified IR, then rebuild in another context. Values
16+
// include precision below float64's significand and NaN payload/signaling bits.
17+
cases := []struct {
18+
name, ir string
19+
words []uint64
20+
}{
21+
{"half", "half 0xH8000", []uint64{0x8000}},
22+
{"bfloat", "bfloat 0xR7FC1", []uint64{0x7fc1}},
23+
{"float_nan", "float 0x7FF82468A0000000", []uint64{0x7fc12345}},
24+
{"double_negative_zero", "double 0x8000000000000000", []uint64{0x8000000000000000}},
25+
{"double_snan", "double 0x7FF0000000000001", []uint64{0x7ff0000000000001}},
26+
{"double_infinity", "double 0x7FF0000000000000", []uint64{0x7ff0000000000000}},
27+
{"double_subnormal", "double 0x0000000000000001", []uint64{1}},
28+
{"fp80_precision", "x86_fp80 0xK3FFF8000000000000001", []uint64{0x8000000000000001, 0x3fff}},
29+
{"fp80_nan", "x86_fp80 0xK7FFFC000000000000123", []uint64{0xc000000000000123, 0x7fff}},
30+
{"fp128_precision", "fp128 0xL00000000000000013FFF000000000000", []uint64{1, 0x3fff000000000000}},
31+
{"fp128_nan", "fp128 0xL00000000000001237FFF800000000000", []uint64{0x123, 0x7fff800000000000}},
32+
{"ppc_fp128_low_only", "ppc_fp128 0xM00000000000000003FF0000000000000", []uint64{0, 0x3ff0000000000000}},
33+
{"ppc_fp128_negative_zero", "ppc_fp128 0xM80000000000000000000000000000000", []uint64{0x8000000000000000, 0}},
34+
{"ppc_fp128_noncanonical", "ppc_fp128 0xM3FF00000000000003FF0000000000000", []uint64{0x3ff0000000000000, 0x3ff0000000000000}},
35+
{"ppc_fp128_nan", "ppc_fp128 0xM7FF80000000001230000000000000000", []uint64{0x7ff8000000000123, 0}},
36+
{"ppc_fp128_precision", "ppc_fp128 0xM3FF00000000000003C90000000000000", []uint64{0x3ff0000000000000, 0x3c90000000000000}},
37+
}
38+
for _, tc := range cases {
39+
t.Run(tc.name, func(t *testing.T) {
40+
srcCtx, dstCtx := NewContext(), NewContext()
41+
defer dstCtx.Dispose()
42+
path := filepath.Join(t.TempDir(), "float.ll")
43+
if err := os.WriteFile(path, []byte("@value = constant "+tc.ir+"\n"), 0600); err != nil {
44+
t.Fatal(err)
45+
}
46+
buf, err := NewMemoryBufferFromFile(path)
47+
if err != nil {
48+
t.Fatal(err)
49+
}
50+
src, err := srcCtx.ParseIR(buf)
51+
if err != nil {
52+
t.Fatal(err)
53+
}
54+
value := src.NamedGlobal("value").Initializer()
55+
words := value.FloatBits()
56+
if !reflect.DeepEqual(words, tc.words) {
57+
t.Fatalf("bits = %x, want %x", words, tc.words)
58+
}
59+
// Reparse just to obtain the identical destination-context type.
60+
buf, err = NewMemoryBufferFromFile(path)
61+
if err != nil {
62+
t.Fatal(err)
63+
}
64+
dst, err := dstCtx.ParseIR(buf)
65+
if err != nil {
66+
t.Fatal(err)
67+
}
68+
defer dst.Dispose()
69+
expected := dst.NamedGlobal("value").Initializer()
70+
cloned := ConstFloatFromBits(expected.Type(), words)
71+
if cloned != expected {
72+
t.Fatalf("rebuilt %s, want %s", cloned, expected)
73+
}
74+
src.Dispose()
75+
srcCtx.Dispose()
76+
if got := cloned.FloatBits(); !reflect.DeepEqual(got, tc.words) {
77+
t.Fatalf("after source disposal: %x", got)
78+
}
79+
words[0] ^= 1
80+
if got := cloned.FloatBits(); !reflect.DeepEqual(got, tc.words) {
81+
t.Fatalf("returned slice aliases constant: %x", got)
82+
}
83+
if err := VerifyModule(dst, ReturnStatusAction); err != nil {
84+
t.Fatal(err)
85+
}
86+
})
87+
}
88+
}
89+
90+
func TestFloatBitsInvalidInput(t *testing.T) {
91+
ctx := NewContext()
92+
defer ctx.Dispose()
93+
for name, call := range map[string]func(){
94+
"nil_type": func() { ConstFloatFromBits(Type{}, []uint64{0}) },
95+
"integer_type": func() { ConstFloatFromBits(ctx.Int64Type(), []uint64{0}) },
96+
"vector_type": func() { ConstFloatFromBits(VectorType(ctx.DoubleType(), 2), []uint64{0, 0}) },
97+
"empty": func() { ConstFloatFromBits(ctx.DoubleType(), nil) },
98+
"too_few": func() { ConstFloatFromBits(ctx.FP128Type(), []uint64{0}) },
99+
"too_many": func() { ConstFloatFromBits(ctx.DoubleType(), []uint64{0, 0}) },
100+
"unused_high_bits": func() { ConstFloatFromBits(ctx.FloatType(), []uint64{1 << 32}) },
101+
"fp80_unused_high_bits": func() { ConstFloatFromBits(ctx.X86FP80Type(), []uint64{0, 1 << 16}) },
102+
"nil_value": func() { Value{}.FloatBits() },
103+
"integer_value": func() { ConstInt(ctx.Int64Type(), 0, false).FloatBits() },
104+
"undef": func() { Undef(ctx.DoubleType()).FloatBits() },
105+
} {
106+
t.Run(name, func(t *testing.T) {
107+
defer func() {
108+
if recover() == nil {
109+
t.Fatal("expected panic")
110+
}
111+
}()
112+
call()
113+
})
114+
}
115+
}

‎ir.go‎

Lines changed: 31 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -927,6 +927,37 @@ func ConstFloatFromString(t Type, str string) (v Value) {
927927
return
928928
}
929929

930+
// ConstFloatFromBits constructs a scalar floating-point constant without rounding
931+
// through float64. words holds the raw representation, least significant word
932+
// first, independently of host byte order. Its length must be ceil(bitWidth/64),
933+
// and unused high bits in the final word must be zero. It panics for an invalid
934+
// type or representation. For ppc_fp128, the first word is the leading double
935+
// and the second word is the trailing double, as in LLVM's APFloat representation.
936+
func ConstFloatFromBits(t Type, words []uint64) (v Value) {
937+
var data *C.uint64_t
938+
if len(words) != 0 {
939+
data = (*C.uint64_t)(unsafe.Pointer(&words[0]))
940+
}
941+
v.C = C.LLVMGoConstFPFromBits(t.C, data, C.uint(len(words)))
942+
if v.IsNil() {
943+
panic("llvm: invalid floating-point type or bit representation")
944+
}
945+
return
946+
}
947+
948+
// FloatBits returns a copy of a scalar floating-point constant's raw APFloat
949+
// representation in least-significant-word-first order. Unused high bits in the
950+
// final word are zero. It panics if v is not a scalar floating-point constant.
951+
func (v Value) FloatBits() []uint64 {
952+
n := C.LLVMGoConstFPGetBits(v.C, nil)
953+
if n == 0 {
954+
panic("llvm: FloatBits requires a scalar floating-point constant")
955+
}
956+
words := make([]uint64, int(n))
957+
C.LLVMGoConstFPGetBits(v.C, (*C.uint64_t)(unsafe.Pointer(&words[0])))
958+
return words
959+
}
960+
930961
func (v Value) ZExtValue() uint64 { return uint64(C.LLVMConstIntGetZExtValue(v.C)) }
931962
func (v Value) SExtValue() int64 { return int64(C.LLVMConstIntGetSExtValue(v.C)) }
932963
func (v Value) DoubleValue() (result float64, inexact bool) {

0 commit comments

Comments
 (0)