Skip to content
Open
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
74 changes: 57 additions & 17 deletions src/schema.jl
Original file line number Diff line number Diff line change
Expand Up @@ -56,8 +56,8 @@ Base.keys(schema::Schema) = keys(schema.schema)
Base.haskey(schema::Schema, key) = haskey(schema.schema, key)

"""
schema([terms::AbstractVector{<:AbstractTerm}, ]data, hints::Dict{Symbol})
schema(term::AbstractTerm, data, hints::Dict{Symbol})
schema([terms::AbstractVector{<:AbstractTerm}, ]data, hints::Dict{Symbol}; statistics=true)
schema(term::AbstractTerm, data, hints::Dict{Symbol}; statistics=true)

Compute all the invariants necessary to fit a model with `terms`. A schema is a dict that
maps `Term`s to their concrete instantiations (either `CategoricalTerm`s or
Expand All @@ -67,6 +67,15 @@ the appropriate term type will be guessed based on the data type from the data c
numeric data is assumed to be continuous, and any non-numeric data is assumed to be
categorical.

By default, creating a `ContinuousTerm` computes the summary statistics (mean, variance,
and extrema) of the corresponding column, which requires several passes over the data.
Pass `statistics=false` to skip this computation and fill those fields with `NaN`
placeholders instead. Use this only when nothing downstream reads the summary statistics
(they are needed, e.g., by packages that use them as default centering or scaling values).
The keyword only affects the terms that StatsModels itself concretizes as continuous
(numeric columns without a hint, and columns with a `ContinuousTerm` hint); terms with
other hints are created exactly as they would be otherwise.

Returns a [`StatsModels.Schema`](@ref), which is a wrapper around a `Dict`
mapping `Term`s to their concrete instantiations (`ContinuousTerm` or
`CategoricalTerm`).
Expand Down Expand Up @@ -114,21 +123,39 @@ julia> sch[term(:y)]
y(continuous)
```
"""
schema(data, hints=Dict{Symbol,Any}()) = schema(columntable(data), hints)
schema(dt::D, hints=Dict{Symbol,Any}()) where {D<:ColumnTable} =
schema(Term.(collect(fieldnames(D))), dt, hints)
schema(ts::AbstractVector{<:AbstractTerm}, data, hints::Dict{Symbol}) =
schema(ts, columntable(data), hints)
schema(data, hints=Dict{Symbol,Any}(); kwargs...) = schema(columntable(data), hints; kwargs...)
schema(dt::D, hints=Dict{Symbol,Any}(); kwargs...) where {D<:ColumnTable} =
schema(Term.(collect(fieldnames(D))), dt, hints; kwargs...)
schema(ts::AbstractVector{<:AbstractTerm}, data, hints::Dict{Symbol}; kwargs...) =
schema(ts, columntable(data), hints; kwargs...)

# handle hints:
schema(ts::AbstractVector{<:AbstractTerm}, dt::ColumnTable,
hints::Dict{Symbol}=Dict{Symbol,Any}()) =
sch = Schema(t=>concrete_term(t, dt, hints) for t in ts)
function schema(ts::AbstractVector{<:AbstractTerm}, dt::ColumnTable,
hints::Dict{Symbol}=Dict{Symbol,Any}(); statistics::Bool=true)
sch = Schema()
for t in ts
# route `statistics=false` only into the paths where the statistics are
# computed by StatsModels itself, so that `concrete_term` methods that
# other packages define for their own hint types never see the keyword
if !statistics && t isa Term
msg = checkcol(dt, t.sym)
msg != "" && throw(ArgumentError(msg))
col = getproperty(dt, t.sym)
hint = get(hints, t.sym, nothing)
if hint === ContinuousTerm || (hint === nothing && col isa AbstractVector{<:Number})
sch.schema[t] = concrete_term(t, col, ContinuousTerm; statistics=false)
continue
end
end
sch.schema[t] = concrete_term(t, dt, hints)
end
return sch
end

schema(f::TermOrTerms, data, hints::Dict{Symbol}) =
schema(filter(needs_schema, terms(f)), data, hints)
schema(f::TermOrTerms, data, hints::Dict{Symbol}; kwargs...) =
schema(filter(needs_schema, terms(f)), data, hints; kwargs...)

schema(f::TermOrTerms, data) = schema(f, data, Dict{Symbol,Any}())
schema(f::TermOrTerms, data; kwargs...) = schema(f, data, Dict{Symbol,Any}(); kwargs...)

"""
concrete_term(t::Term, data[, hint])
Expand All @@ -146,6 +173,11 @@ If no hint is provided (or `hint==nothing`), the `eltype` of the data is used:
`Number`s are assumed to be continuous, and all others are assumed to be
categorical.

The `ContinuousTerm`-hint method additionally accepts a `statistics::Bool=true`
keyword: `concrete_term(t, xs, ContinuousTerm, statistics=false)` skips computing
the summary statistics (mean, variance, and extrema) stored in the term, filling
those fields with `NaN` placeholders instead (see [`schema`](@ref)).

# Example

```jldoctest
Expand Down Expand Up @@ -198,10 +230,18 @@ concrete_term(t::Term, x, hint::AbstractTerm) = hint
concrete_term(t, d, hint) = t

concrete_term(t::Term, xs::AbstractVector{<:Number}, ::Nothing) = concrete_term(t, xs, ContinuousTerm)
function concrete_term(t::Term, xs::AbstractVector, ::Type{ContinuousTerm})
μ, σ2 = StatsBase.mean_and_var(xs)
min, max = extrema(xs)
ContinuousTerm(t.sym, promote(μ, σ2, min, max)...)
function concrete_term(t::Term, xs::AbstractVector, ::Type{ContinuousTerm};
statistics::Bool=true)
if statistics
μ, σ2 = StatsBase.mean_and_var(xs)
min, max = extrema(xs)
return ContinuousTerm(t.sym, promote(μ, σ2, min, max)...)
else
# keep the field type the statistics would have had, without the O(n) passes
E = eltype(xs)
nan = E <: Number ? convert(float(E), NaN) : NaN
return ContinuousTerm(t.sym, nan, nan, nan, nan)
end
end
# default contrasts: dummy coding
concrete_term(t::Term, xs::AbstractVector, ::Nothing) = concrete_term(t, xs, CategoricalTerm)
Expand Down
3 changes: 3 additions & 0 deletions src/terms.jl
Original file line number Diff line number Diff line change
Expand Up @@ -196,6 +196,9 @@ Represents a continuous variable, with a name and summary statistics.
* `var::T`: Variance
* `min::T`: Minimum value
* `max::T`: Maximum value

The summary statistics are `NaN` placeholders when the term was created with
`statistics=false` (see [`schema`](@ref)).
"""
struct ContinuousTerm{T} <: AbstractTerm
sym::Symbol
Expand Down
35 changes: 35 additions & 0 deletions test/schema.jl
Original file line number Diff line number Diff line change
Expand Up @@ -70,6 +70,41 @@

end

@testset "statistics=false" begin
f = @formula(y ~ 1 + a + b + c)
d = (y = rand(10), a = rand(10), b = Float32.(1:10),
c = repeat(["u","v"], 5))

sch = schema(f, d, statistics=false)
t = sch[term(:a)]
@test t isa ContinuousTerm
@test isnan(t.mean) && isnan(t.var) && isnan(t.min) && isnan(t.max)
# the statistics keep the type they would have had
@test sch[term(:b)] isa ContinuousTerm{Float32}
# categorical terms are unaffected
c1, c2 = sch[term(:c)], schema(f, d)[term(:c)]
@test c1 isa CategoricalTerm && c1.sym == c2.sym && c1.contrasts == c2.contrasts

# skipping statistics changes nothing else downstream
ff = apply_schema(f, sch)
@test modelmatrix(ff.rhs, d) == modelmatrix(apply_schema(f, schema(f, d)).rhs, d)

# hints still work, and never see the keyword
sch1 = schema(f, d, Dict(:a => CategoricalTerm), statistics=false)
@test sch1[term(:a)] isa CategoricalTerm{DummyCoding}
sch2 = schema(f, d, Dict(:a => EffectsCoding()), statistics=false)
@test sch2[term(:a)] isa CategoricalTerm{EffectsCoding}
# a ContinuousTerm hint skips the statistics too
sch3 = schema(f, d, Dict(:a => ContinuousTerm), statistics=false)
@test isnan(sch3[term(:a)].mean)
# an AbstractTerm hint is included as is
hint = schema(f, d)[term(:a)]
@test schema(f, d, Dict(:a => hint), statistics=false)[term(:a)] === hint

@test isnan(concrete_term(term(:a), [1, 2, 3], ContinuousTerm, statistics=false).mean)
@test concrete_term(term(:a), [1, 2, 3], ContinuousTerm, statistics=true).mean == 2.0
end

@testset "nice errors" begin
d = (yyy = rand(10), aaa = rand(10), bbb = repeat([:a, :b], 5))
@test_throws ArgumentError("There isn't a variable called 'aa' in your data; the nearest names appear to be: aaa") concrete_term(Term(:aa), d, nothing)
Expand Down
Loading